vortex_sequence/compute/
list_contains.rs1use vortex_array::ArrayRef;
5use vortex_array::ArrayView;
6use vortex_array::IntoArray;
7use vortex_array::arrays::BoolArray;
8use vortex_array::arrays::ConstantArray;
9use vortex_array::scalar::Scalar;
10use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce;
11use vortex_error::VortexExpect;
12use vortex_error::VortexResult;
13
14use crate::array::Sequence;
15use crate::compute::compare::Intersection;
16use crate::compute::compare::find_intersection;
17
18impl ListContainsElementReduce for Sequence {
19 fn list_contains(
20 list: &ArrayRef,
21 element: ArrayView<'_, Self>,
22 ) -> VortexResult<Option<ArrayRef>> {
23 let Some(list_scalar) = list.as_constant() else {
24 return Ok(None);
25 };
26
27 let list_elements = list_scalar
28 .as_list()
29 .elements()
30 .vortex_expect("non-null element (checked in entry)");
31
32 let nullability = list.dtype().nullability() | element.dtype().nullability();
33
34 let mut set_indices: Vec<usize> = Vec::new();
35 for intercept in list_elements.iter() {
36 let Some(intercept) = intercept.as_primitive().pvalue() else {
37 continue;
38 };
39 match find_intersection(
40 element.base(),
41 element.multiplier(),
42 element.len(),
43 intercept,
44 ) {
45 None | Some(Intersection::None) => {}
47 Some(Intersection::At(idx)) => set_indices.push(idx),
48 Some(Intersection::All) => {
49 return Ok(Some(
50 ConstantArray::new(Scalar::bool(true, nullability), element.len())
51 .into_array(),
52 ));
53 }
54 }
55 }
56
57 Ok(Some(
58 BoolArray::from_indices(element.len(), set_indices, nullability.into()).into_array(),
59 ))
60 }
61}
62
63#[cfg(test)]
64mod tests {
65 use std::sync::Arc;
66 use std::sync::LazyLock;
67
68 use vortex_array::IntoArray;
69 use vortex_array::VortexSessionExecute;
70 use vortex_array::arrays::BoolArray;
71 use vortex_array::assert_arrays_eq;
72 use vortex_array::dtype::Nullability;
73 use vortex_array::dtype::PType::I32;
74 use vortex_array::expr::list_contains;
75 use vortex_array::expr::lit;
76 use vortex_array::expr::root;
77 use vortex_array::scalar::Scalar;
78 use vortex_session::VortexSession;
79
80 use crate::Sequence;
81
82 static SESSION: LazyLock<VortexSession> = LazyLock::new(|| {
83 let session = vortex_array::array_session();
84 crate::initialize(&session);
85 session
86 });
87
88 #[test]
89 fn test_list_contains_seq() {
90 let list_scalar = Scalar::list(
91 Arc::new(I32.into()),
92 vec![1.into(), 3.into()],
93 Nullability::Nullable,
94 );
95
96 {
97 let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3)
101 .unwrap()
102 .into_array();
103
104 let expr = list_contains(lit(list_scalar.clone()), root());
105 let result = array.into_array().apply(&expr).unwrap();
106 let expected = BoolArray::from_iter([Some(true), Some(false), Some(true)]);
107 assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
108 }
109
110 {
111 let array = Sequence::try_new_typed(1, 2, Nullability::NonNullable, 3)
115 .unwrap()
116 .into_array();
117
118 let expr = list_contains(lit(list_scalar), root());
119 let result = array.into_array().apply(&expr).unwrap();
120 let expected = BoolArray::from_iter([Some(true), Some(true), Some(false)]);
121 assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
122 }
123 }
124
125 #[test]
126 fn test_list_contains_constant_sequence() {
127 let list_scalar = Scalar::list(
128 Arc::new(I32.into()),
129 vec![7.into(), 42.into()],
130 Nullability::Nullable,
131 );
132
133 let array = Sequence::try_new_typed(42i32, 0, Nullability::NonNullable, 3)
134 .unwrap()
135 .into_array();
136
137 let expr = list_contains(lit(list_scalar), root());
138 let result = array.apply(&expr).unwrap();
139 let expected = BoolArray::from_iter([Some(true), Some(true), Some(true)]);
140 assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx());
141 }
142}