diff --git a/crates/paimon/src/arrow/residual.rs b/crates/paimon/src/arrow/residual.rs index 60cfe401..56b77050 100644 --- a/crates/paimon/src/arrow/residual.rs +++ b/crates/paimon/src/arrow/residual.rs @@ -540,6 +540,10 @@ fn evaluate_set_membership_predicate( }); } + if let Some(mask) = set_membership_hash_mask(array, op, literals) { + return Ok(mask); + } + let mut combined = match op { PredicateOperator::In => BooleanArray::from(vec![false; array.len()]), PredicateOperator::NotIn => { @@ -566,6 +570,132 @@ fn evaluate_set_membership_predicate( Ok(combined) } +/// One-pass hash-set evaluation of `In`/`NotIn` for byte-like, string, and +/// integer columns. The general path above OR-combines one comparison kernel +/// per literal — O(rows × literals) — which turns a large pushed-down literal +/// set (an engine probing a batch of keys as `In`) quadratic. Returns `None` +/// for column/literal shapes outside the fast path, or when any literal would +/// not convert for the column, so the general path keeps its exact semantics — +/// including its error behavior for unconvertible literals. +fn set_membership_hash_mask( + array: &ArrayRef, + op: PredicateOperator, + literals: &[Datum], +) -> Option { + use arrow_schema::DataType as ArrowType; + use std::collections::HashSet; + + let keep = matches!(op, PredicateOperator::In); + + fn byte_mask<'a>( + values: impl Iterator>, + literals: &[Datum], + keep: bool, + ) -> Option { + let set = literals + .iter() + .map(|literal| match literal { + Datum::Bytes(bytes) => Some(bytes.as_slice()), + _ => None, + }) + .collect::>>()?; + Some( + values + .map(|value| Some(value.is_some_and(|v| set.contains(v) == keep))) + .collect(), + ) + } + + fn str_mask<'a>( + values: impl Iterator>, + literals: &[Datum], + keep: bool, + ) -> Option { + let set = literals + .iter() + .map(|literal| match literal { + Datum::String(value) => Some(value.as_str()), + _ => None, + }) + .collect::>>()?; + Some( + values + .map(|value| Some(value.is_some_and(|v| set.contains(v) == keep))) + .collect(), + ) + } + + fn int_mask(array: &ArrayRef, literals: &[Datum], keep: bool) -> Option + where + T: arrow_array::types::ArrowPrimitiveType, + T::Native: TryFrom + std::hash::Hash + Eq, + { + let array = array + .as_any() + .downcast_ref::>()?; + let set = literals + .iter() + .map(|literal| integer_literal(literal).and_then(|v| T::Native::try_from(v).ok())) + .collect::>>()?; + Some( + array + .iter() + .map(|value| Some(value.is_some_and(|v| set.contains(&v) == keep))) + .collect(), + ) + } + + match array.data_type() { + ArrowType::Binary => byte_mask( + array.as_any().downcast_ref::()?.iter(), + literals, + keep, + ), + ArrowType::LargeBinary => byte_mask( + array + .as_any() + .downcast_ref::()? + .iter(), + literals, + keep, + ), + ArrowType::BinaryView => byte_mask( + array + .as_any() + .downcast_ref::()? + .iter(), + literals, + keep, + ), + ArrowType::Utf8 => str_mask( + array.as_any().downcast_ref::()?.iter(), + literals, + keep, + ), + ArrowType::LargeUtf8 => str_mask( + array + .as_any() + .downcast_ref::()? + .iter(), + literals, + keep, + ), + ArrowType::Utf8View => str_mask( + array + .as_any() + .downcast_ref::()? + .iter(), + literals, + keep, + ), + ArrowType::Int8 => int_mask::(array, literals, keep), + ArrowType::Int16 => int_mask::(array, literals, keep), + ArrowType::Int32 => int_mask::(array, literals, keep), + ArrowType::Int64 => int_mask::(array, literals, keep), + _ => None, + } +} + fn evaluate_column_predicate( column: &ArrayRef, scalar: &Scalar, @@ -985,6 +1115,87 @@ mod tests { (0..col.len()).map(|i| col.value(i).to_string()).collect() } + #[test] + fn test_in_hash_path_filters_exactly_with_nulls() { + let f = int_field(0, "age"); + let b = int_batch("age", vec![Some(10), None, Some(20), Some(40), Some(50)]); + let pred = leaf( + 0, + DataType::Int(IntType::new()), + PredicateOperator::In, + vec![Datum::Int(20), Datum::Int(40), Datum::Int(999)], + ); + let fp = file_predicates(vec![pred], vec![f.clone()]); + let out = filter_record_batch_by_predicates(b, &fp, &[f]).unwrap(); + assert_eq!(int_values(&out), vec![20, 40]); + } + + #[test] + fn test_not_in_hash_path_excludes_nulls() { + let f = str_field(0, "name"); + let b = str_batch( + "name", + vec![Some("apple"), None, Some("banana"), Some("cherry")], + ); + let pred = leaf( + 0, + DataType::VarChar(VarCharType::string_type()), + PredicateOperator::NotIn, + vec![Datum::String("banana".to_string())], + ); + let fp = file_predicates(vec![pred], vec![f.clone()]); + let out = filter_record_batch_by_predicates(b, &fp, &[f]).unwrap(); + assert_eq!(str_values(&out), vec!["apple", "cherry"]); + } + + #[test] + fn test_in_hash_path_on_binary_column() { + use crate::spec::VarBinaryType; + let array: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(b"aa".as_slice()), + None, + Some(b"bb".as_slice()), + Some(b"cc".as_slice()), + ])); + let data_type = + DataType::VarBinary(VarBinaryType::try_new(true, VarBinaryType::MAX_LENGTH).unwrap()); + let literals = vec![Datum::Bytes(b"bb".to_vec()), Datum::Bytes(b"zz".to_vec())]; + let mask = + evaluate_set_membership_predicate(&array, &data_type, PredicateOperator::In, &literals) + .unwrap(); + assert_eq!( + mask.iter().collect::>(), + vec![Some(false), Some(false), Some(true), Some(false)] + ); + let mask = evaluate_set_membership_predicate( + &array, + &data_type, + PredicateOperator::NotIn, + &literals, + ) + .unwrap(); + assert_eq!( + mask.iter().collect::>(), + vec![Some(true), Some(false), Some(false), Some(true)] + ); + } + + #[test] + fn test_in_unconvertible_literal_still_errors() { + // An out-of-range literal must keep the general path's error behavior: + // the hash path declines the literal set and the per-literal loop + // raises the unconvertible-literal error it always has. + use crate::spec::TinyIntType; + let array: ArrayRef = Arc::new(arrow_array::Int8Array::from(vec![Some(1i8), Some(2)])); + let result = evaluate_set_membership_predicate( + &array, + &DataType::TinyInt(TinyIntType::new()), + PredicateOperator::In, + &[Datum::Long(300)], + ); + assert!(result.is_err()); + } + #[test] fn test_gt_filters_exactly() { let f = int_field(0, "age");