Skip to main content

vortex_sequence/compute/
list_contains.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use 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                // Non-integer elements do not match the sequence.
46                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            // [1, 3] in  1
98            //            2
99            //            3
100            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            // [1, 3] in  1
112            //            3
113            //            5
114            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}