Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

85 changes: 74 additions & 11 deletions encodings/sequence/src/compute/list_contains.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,9 @@ use vortex_array::arrays::BoolArray;
use vortex_array::arrays::ConstantArray;
use vortex_array::scalar::Scalar;
use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce;
use vortex_error::VortexExpect;
use vortex_array::scalar_fn::fns::list_contains::ListContainsOptions;
use vortex_array::validity::Validity;
use vortex_buffer::BitBufferMut;
use vortex_error::VortexResult;

use crate::array::Sequence;
Expand All @@ -19,21 +21,24 @@ impl ListContainsElementReduce for Sequence {
fn list_contains(
list: &ArrayRef,
element: ArrayView<'_, Self>,
options: &ListContainsOptions,
) -> VortexResult<Option<ArrayRef>> {
let Some(list_scalar) = list.as_constant() else {
return Ok(None);
};

let list_elements = list_scalar
.as_list()
.elements()
.vortex_expect("non-null element (checked in entry)");
// A null list falls back to the generic path, which yields all-null.
let Some(list_elements) = list_scalar.as_list().elements() else {
return Ok(None);
};

let nullability = list.dtype().nullability() | element.dtype().nullability();
let nullability = options.result_nullability(list.dtype(), element.dtype());

let mut set_indices: Vec<usize> = Vec::new();
let mut matches = BitBufferMut::new_unset(element.len());
let mut has_null = false;
for intercept in list_elements.iter() {
let Some(intercept) = intercept.as_primitive().pvalue() else {
has_null = true;
continue;
};
match find_intersection(
Expand All @@ -44,7 +49,7 @@ impl ListContainsElementReduce for Sequence {
) {
// Non-integer elements do not match the sequence.
None | Some(Intersection::None) => {}
Some(Intersection::At(idx)) => set_indices.push(idx),
Some(Intersection::At(idx)) => matches.set(idx),
Some(Intersection::All) => {
return Ok(Some(
ConstantArray::new(Scalar::bool(true, nullability), element.len())
Expand All @@ -54,9 +59,16 @@ impl ListContainsElementReduce for Sequence {
}
}

Ok(Some(
BoolArray::from_indices(element.len(), set_indices, nullability.into()).into_array(),
))
let matches = matches.freeze();

// Under SQL null semantics a null element makes every non-match unknown.
let validity = if options.sql_null_semantics && has_null {
Validity::from_bit_buffer(matches.clone(), nullability)
} else {
nullability.into()
};

Ok(Some(BoolArray::new(matches, validity).into_array()))
}
}

Expand All @@ -69,12 +81,15 @@ mod tests {
use vortex_array::VortexSessionExecute;
use vortex_array::arrays::BoolArray;
use vortex_array::assert_arrays_eq;
use vortex_array::dtype::DType;
use vortex_array::dtype::Nullability;
use vortex_array::dtype::PType::I32;
use vortex_array::expr::list_contains;
use vortex_array::expr::list_contains_opts;
use vortex_array::expr::lit;
use vortex_array::expr::root;
use vortex_array::scalar::Scalar;
use vortex_array::scalar_fn::fns::list_contains::ListContainsOptions;
use vortex_session::VortexSession;

use crate::Sequence;
Expand Down Expand Up @@ -122,6 +137,54 @@ mod tests {
}
}

#[test]
fn test_list_contains_null_list() {
// A null list yields null for every row. The reduce rule used to panic on it.
let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3)
.unwrap()
.into_array();

let null_list = Scalar::null(DType::List(Arc::new(I32.into()), Nullability::Nullable));
let expr = list_contains(lit(null_list), root());
let result = array.apply(&expr).unwrap();
let expected = BoolArray::from_iter([None::<bool>, None, None]);
assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
}

#[test]
fn test_list_contains_null_element_semantics() {
// By default a null element matches nothing. Under SQL null semantics it makes every
// non-match null.
let element = DType::Primitive(I32, Nullability::Nullable);
let set = Scalar::list(
Arc::new(element.clone()),
vec![
Scalar::primitive(1i32, Nullability::Nullable),
Scalar::null(element),
],
Nullability::NonNullable,
);
let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3)
.unwrap()
.into_array();

let result = array
.clone()
.apply(&list_contains(lit(set.clone()), root()))
.unwrap();
let expected = BoolArray::from_iter([true, false, false]);
assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());

let sql = ListContainsOptions {
sql_null_semantics: true,
};
let result = array
.apply(&list_contains_opts(lit(set), root(), sql))
.unwrap();
let expected = BoolArray::from_iter([Some(true), None, None]);
assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
}

#[test]
fn test_list_contains_constant_sequence() {
let list_scalar = Scalar::list(
Expand Down
4 changes: 4 additions & 0 deletions vortex-array/proto/expr.proto
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,10 @@ message LikeOpts {
bool case_insensitive = 2;
}

message ListContainsOpts {
bool sql_null_semantics = 1;
}

message CastOpts {
vortex.dtype.DType target = 1;
}
Expand Down
3 changes: 2 additions & 1 deletion vortex-array/src/builtins.rs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ use crate::scalar_fn::fns::get_item::GetItem;
use crate::scalar_fn::fns::is_not_null::IsNotNull;
use crate::scalar_fn::fns::is_null::IsNull;
use crate::scalar_fn::fns::list_contains::ListContains;
use crate::scalar_fn::fns::list_contains::ListContainsOptions;
use crate::scalar_fn::fns::mask::Mask;
use crate::scalar_fn::fns::not::Not;
use crate::scalar_fn::fns::operators::Operator;
Expand Down Expand Up @@ -103,7 +104,7 @@ impl ExprBuiltins for Expression {
}

fn list_contains(&self, value: Expression) -> VortexResult<Expression> {
ListContains.try_new_expr(EmptyOptions, [self.clone(), value])
ListContains.try_new_expr(ListContainsOptions::default(), [self.clone(), value])
}

fn zip(&self, if_true: Expression, if_false: Expression) -> VortexResult<Expression> {
Expand Down
55 changes: 49 additions & 6 deletions vortex-array/src/expr/exprs.rs
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ use crate::scalar_fn::fns::is_null::IsNull;
use crate::scalar_fn::fns::like::Like;
use crate::scalar_fn::fns::like::LikeOptions;
use crate::scalar_fn::fns::list_contains::ListContains;
use crate::scalar_fn::fns::list_contains::ListContainsOptions;
use crate::scalar_fn::fns::list_length::ListLength;
use crate::scalar_fn::fns::list_sum::ListSum;
use crate::scalar_fn::fns::literal::Literal;
Expand Down Expand Up @@ -1120,23 +1121,64 @@ pub fn bound_dynamic(

// ---- ListContains ----

/// Creates an expression that checks if a value is contained in a list.
/// Creates an expression that checks whether a list contains a value.
///
/// Returns a boolean array indicating whether the value appears in each list.
/// A null element never matches anything: a needle that matches no element yields `false`. For
/// SQL's three-valued `IN`, where a null element makes that `null`, use [`in_list`].
///
/// ```rust
/// # use vortex_array::expr::{list_contains, lit, root};
/// let expr = list_contains(root(), lit(42));
/// ```
pub fn list_contains(list: Expression, value: Expression) -> Expression {
ListContains.new_expr(EmptyOptions, [list, value])
list_contains_opts(list, value, ListContainsOptions::default())
}

/// Creates a bound expression that checks if a value is contained in a list.
/// Creates an expression that checks whether a list contains a value, with explicit
/// [`ListContainsOptions`].
pub fn list_contains_opts(
list: Expression,
value: Expression,
options: ListContainsOptions,
) -> Expression {
ListContains.new_expr(options, [list, value])
}

/// Creates `needle IN (list)` with SQL null semantics: a null element is an unknown value, so a
/// needle that matches no element yields `null` rather than `false` when the list holds one, and
/// `not(in_list(..))` never admits such a row.
///
/// ```rust
/// # use vortex_array::expr::{in_list, lit, root};
/// let expr = in_list(root(), lit(vec![1, 2, 3]));
/// ```
pub fn in_list(needle: Expression, list: Expression) -> Expression {
list_contains_opts(
list,
needle,
ListContainsOptions {
sql_null_semantics: true,
},
)
}

/// Creates a bound expression that checks whether a list contains a value.
pub fn bound_list_contains(list: BoundExpression, value: BoundExpression) -> BoundExpression {
ListContains
.try_new_bound_expr(EmptyOptions, [list, value])
.vortex_expect("list-contains expressions require a compatible list and value dtype")
.try_new_bound_expr(ListContainsOptions::default(), [list, value])
.vortex_expect("list-contains expressions require a list child")
}

/// Creates a bound `needle IN (list)` with SQL null semantics. See [`in_list`].
pub fn bound_in_list(needle: BoundExpression, list: BoundExpression) -> BoundExpression {
ListContains
.try_new_bound_expr(
ListContainsOptions {
sql_null_semantics: true,
},
[list, needle],
)
.vortex_expect("list-contains expressions require a list child")
}

// ---- ByteLength ----
Expand Down Expand Up @@ -1265,6 +1307,7 @@ pub mod bound {
pub use super::bound_gt as gt;
pub use super::bound_gt_eq as gt_eq;
pub use super::bound_ilike as ilike;
pub use super::bound_in_list as in_list;
pub use super::bound_is_nan as is_nan;
pub use super::bound_is_not_null as is_not_null;
pub use super::bound_is_null as is_null;
Expand Down
2 changes: 2 additions & 0 deletions vortex-array/src/expr/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -79,12 +79,14 @@ pub use exprs::get_item;
pub use exprs::gt;
pub use exprs::gt_eq;
pub use exprs::ilike;
pub use exprs::in_list;
pub use exprs::is_nan;
pub use exprs::is_not_null;
pub use exprs::is_null;
pub use exprs::is_root;
pub use exprs::like;
pub use exprs::list_contains;
pub use exprs::list_contains_opts;
pub use exprs::list_length;
pub use exprs::list_sum;
pub use exprs::list_sum_opts;
Expand Down
5 changes: 5 additions & 0 deletions vortex-array/src/proto/generated/vortex.expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,11 @@ pub struct LikeOpts {
#[prost(bool, tag = "2")]
pub case_insensitive: bool,
}
#[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)]
pub struct ListContainsOpts {
#[prost(bool, tag = "1")]
pub sql_null_semantics: bool,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct CastOpts {
#[prost(message, optional, tag = "1")]
Expand Down
7 changes: 5 additions & 2 deletions vortex-array/src/scalar_fn/fns/list_contains/kernel.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ use crate::arrays::scalar_fn::ScalarFnArrayView;
use crate::kernel::ExecuteParentKernel;
use crate::optimizer::rules::ArrayParentReduceRule;
use crate::scalar_fn::fns::list_contains::ListContains as ListContainsExpr;
use crate::scalar_fn::fns::list_contains::ListContainsOptions;

/// Check list-contains without reading buffers (metadata-only).
///
Expand All @@ -30,6 +31,7 @@ pub trait ListContainsElementReduce: VTable {
fn list_contains(
list: &ArrayRef,
element: ArrayView<'_, Self>,
options: &ListContainsOptions,
) -> VortexResult<Option<ArrayRef>>;
}

Expand All @@ -42,6 +44,7 @@ pub trait ListContainsElementKernel: VTable {
fn list_contains(
list: &ArrayRef,
element: ArrayView<'_, Self>,
options: &ListContainsOptions,
ctx: &mut ExecutionCtx,
) -> VortexResult<Option<ArrayRef>>;
}
Expand Down Expand Up @@ -70,7 +73,7 @@ where
.as_opt::<ScalarFn>()
.vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray");
let list = scalar_fn_array.get_child(0);
<V as ListContainsElementReduce>::list_contains(list, array)
<V as ListContainsElementReduce>::list_contains(list, array, parent.options)
}
}

Expand Down Expand Up @@ -99,6 +102,6 @@ where
.as_opt::<ScalarFn>()
.vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray");
let list = scalar_fn_array.get_child(0);
<V as ListContainsElementKernel>::list_contains(list, array, ctx)
<V as ListContainsElementKernel>::list_contains(list, array, parent.options, ctx)
}
}
Loading
Loading