Skip to main content

datafusion_functions_nested/
array_filter.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! [`datafusion_expr::HigherOrderUDF`] definitions for array_filter function.
19
20use arrow::{
21    array::{
22        Array, ArrayRef, AsArray, BooleanArray, LargeListArray, ListArray,
23        OffsetSizeTrait, new_empty_array,
24    },
25    buffer::{OffsetBuffer, ScalarBuffer},
26    compute::filter as arrow_filter,
27    datatypes::{DataType, Field, FieldRef},
28};
29use datafusion_common::{Result, ScalarValue, exec_err};
30use datafusion_expr::{
31    ColumnarValue, Documentation, HigherOrderFunctionArgs, HigherOrderReturnFieldArgs,
32    HigherOrderSignature, HigherOrderUDFImpl, LambdaParametersProgress, ValueOrLambda,
33    Volatility,
34};
35use datafusion_macros::user_doc;
36use std::sync::Arc;
37
38use crate::lambda_utils::{
39    SingleListLambdaResult, coerce_single_list_arg, evaluate_single_list_predicate,
40    single_list_lambda_parameters, value_lambda_pair,
41};
42
43make_higher_order_function_expr_and_func!(
44    ArrayFilter,
45    array_filter,
46    array lambda,
47    "filters the values of an array using a boolean lambda",
48    array_filter_higher_order_function
49);
50
51#[user_doc(
52    doc_section(label = "Array Functions"),
53    description = "filters the values of an array using a boolean lambda",
54    syntax_example = "array_filter(array, x -> x > 2)",
55    sql_example = r#"```sql
56> select array_filter([1, 2, 3, 4, 5], x -> x > 2);
57+--------------------------------------------+
58| array_filter([1, 2, 3, 4, 5], x -> x > 2) |
59+--------------------------------------------+
60| [3, 4, 5]                                  |
61+--------------------------------------------+
62```"#,
63    argument(
64        name = "array",
65        description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
66    ),
67    argument(
68        name = "lambda",
69        description = "Lambda that returns a boolean. Elements for which the lambda returns true are kept."
70    )
71)]
72#[derive(Debug, PartialEq, Eq, Hash)]
73pub struct ArrayFilter {
74    signature: HigherOrderSignature,
75    aliases: Vec<String>,
76}
77
78impl Default for ArrayFilter {
79    fn default() -> Self {
80        Self::new()
81    }
82}
83
84impl ArrayFilter {
85    pub fn new() -> Self {
86        Self {
87            signature: HigherOrderSignature::exact(
88                vec![ValueOrLambda::Value(()), ValueOrLambda::Lambda(())],
89                Volatility::Immutable,
90            ),
91            aliases: vec![String::from("list_filter")],
92        }
93    }
94}
95
96impl HigherOrderUDFImpl for ArrayFilter {
97    fn name(&self) -> &str {
98        "array_filter"
99    }
100
101    fn aliases(&self) -> &[String] {
102        &self.aliases
103    }
104
105    fn signature(&self) -> &HigherOrderSignature {
106        &self.signature
107    }
108
109    fn lambda_parameters(
110        &self,
111        _step: usize,
112        fields: &[ValueOrLambda<FieldRef, Option<FieldRef>>],
113    ) -> Result<LambdaParametersProgress> {
114        single_list_lambda_parameters(self.name(), fields)
115    }
116
117    fn return_field_from_args(
118        &self,
119        args: HigherOrderReturnFieldArgs,
120    ) -> Result<Arc<Field>> {
121        let (list, _lambda) = value_lambda_pair(self.name(), args.arg_fields)?;
122        Ok(Arc::new(Field::new(
123            "",
124            list.data_type().clone(),
125            list.is_nullable(),
126        )))
127    }
128
129    fn invoke_with_args(&self, args: HigherOrderFunctionArgs) -> Result<ColumnarValue> {
130        let evaluated = match evaluate_single_list_predicate(self.name(), &args)? {
131            SingleListLambdaResult::EarlyReturn(v) => return Ok(v),
132            SingleListLambdaResult::Ready(v) => v,
133        };
134
135        let field = match args.return_field.data_type() {
136            DataType::List(field) | DataType::LargeList(field) => Arc::clone(field),
137            _ => {
138                return exec_err!(
139                    "{} expected return_field to be a list, got {}",
140                    self.name(),
141                    args.return_field
142                );
143            }
144        };
145
146        // Scalar predicate short-circuit: x -> true or x -> false/null
147        if let ColumnarValue::Scalar(ScalarValue::Boolean(b)) =
148            &evaluated.evaluated_result
149        {
150            return match b {
151                Some(true) => Ok(ColumnarValue::Array(evaluated.original_list)),
152                _ => Ok(ColumnarValue::Array(empty_filtered_list(
153                    &evaluated.original_list,
154                    field,
155                )?)),
156            };
157        }
158
159        let predicate = evaluated.boolean_predicate(self.name())?;
160
161        // ListView and LargeListView are coerced to List/LargeList by coerce_value_types.
162        let filtered_list = match evaluated.original_list.data_type() {
163            DataType::List(_) => {
164                let (filtered_values, new_offsets) = filter_list_values(
165                    &evaluated.flattened_values,
166                    &predicate,
167                    &evaluated.adjusted_offsets::<i32>(),
168                )?;
169                Arc::new(ListArray::new(
170                    field,
171                    new_offsets,
172                    filtered_values,
173                    evaluated.nulls().cloned(),
174                )) as ArrayRef
175            }
176            DataType::LargeList(_) => {
177                let (filtered_values, new_offsets) = filter_list_values(
178                    &evaluated.flattened_values,
179                    &predicate,
180                    &evaluated.adjusted_offsets::<i64>(),
181                )?;
182                Arc::new(LargeListArray::new(
183                    field,
184                    new_offsets,
185                    filtered_values,
186                    evaluated.nulls().cloned(),
187                ))
188            }
189            other => exec_err!("expected list, got {other}")?,
190        };
191
192        Ok(ColumnarValue::Array(filtered_list))
193    }
194
195    fn coerce_value_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
196        coerce_single_list_arg(self.name(), arg_types)
197    }
198
199    fn documentation(&self) -> Option<&Documentation> {
200        self.doc()
201    }
202}
203
204/// Returns a list array with every non-null sublist emptied, preserving the null buffer.
205/// Used for the `x -> false` / `x -> null` scalar predicate short-circuit.
206fn empty_filtered_list(list_array: &ArrayRef, field: FieldRef) -> Result<ArrayRef> {
207    let n = list_array.len();
208    let empty_values = new_empty_array(field.data_type());
209    Ok(match list_array.data_type() {
210        DataType::List(_) => {
211            let list = list_array.as_list::<i32>();
212            Arc::new(ListArray::new(
213                field,
214                OffsetBuffer::new(ScalarBuffer::from(vec![0i32; n + 1])),
215                empty_values,
216                list.nulls().cloned(),
217            ))
218        }
219        DataType::LargeList(_) => {
220            let list = list_array.as_list::<i64>();
221            Arc::new(LargeListArray::new(
222                field,
223                OffsetBuffer::new(ScalarBuffer::from(vec![0i64; n + 1])),
224                empty_values,
225                list.nulls().cloned(),
226            ))
227        }
228        other => return exec_err!("expected list, got {other}"),
229    })
230}
231
232/// Filters flat list values using a boolean predicate, returning filtered values and
233/// recomputed per-sublist offsets. Null predicate values are treated as false.
234fn filter_list_values<O: OffsetSizeTrait>(
235    values: &ArrayRef,
236    predicate: &BooleanArray,
237    offsets: &OffsetBuffer<O>,
238) -> Result<(ArrayRef, OffsetBuffer<O>)> {
239    let num_sublists = offsets.len().saturating_sub(1);
240    let has_nulls = predicate.null_count() > 0;
241    let new_offsets = OffsetBuffer::<O>::from_lengths((0..num_sublists).map(|i| {
242        let start = offsets[i].as_usize();
243        let end = offsets[i + 1].as_usize();
244        if has_nulls {
245            (start..end)
246                .filter(|&j| predicate.is_valid(j) && predicate.value(j))
247                .count()
248        } else {
249            predicate
250                .values()
251                .slice(start, end - start)
252                .count_set_bits()
253        }
254    }));
255
256    if new_offsets.last() == offsets.last() {
257        return Ok((Arc::clone(values), offsets.clone()));
258    }
259
260    // arrow_filter treats null predicate values as false
261    let filtered_values = arrow_filter(values.as_ref(), predicate)?;
262    Ok((filtered_values, new_offsets))
263}
264
265#[cfg(test)]
266mod tests {
267    use arrow::{
268        array::{Array, AsArray},
269        buffer::{NullBuffer, OffsetBuffer},
270    };
271
272    use arrow::array::Int32Array;
273
274    use crate::array_filter::array_filter_higher_order_function;
275    use crate::lambda_utils::test_utils::{
276        create_i32_large_list, create_i32_list, eval_hof_on_i32_list,
277        eval_hof_on_i32_list_with_outer, v,
278    };
279    use datafusion_expr::{col, lit};
280
281    fn keep_greater_than_two(
282        list: impl Array + Clone + 'static,
283    ) -> datafusion_common::Result<arrow::array::ArrayRef> {
284        eval_hof_on_i32_list(
285            array_filter_higher_order_function(),
286            list,
287            v().gt(lit(2i32)),
288        )
289    }
290
291    #[test]
292    fn filter_basic() {
293        let list = create_i32_list(
294            vec![1, 2, 3, 4, 5],
295            OffsetBuffer::<i32>::from_lengths(vec![5]),
296            None,
297        );
298
299        let res = keep_greater_than_two(list).unwrap();
300        let actual = res.as_list::<i32>();
301
302        let expected = create_i32_list(
303            vec![3, 4, 5],
304            OffsetBuffer::<i32>::from_lengths(vec![3]),
305            None,
306        );
307
308        assert_eq!(actual, &expected);
309    }
310
311    #[test]
312    fn filter_multiple_sublists() {
313        let list = create_i32_list(
314            vec![1, 5, 2, 4, 3],
315            OffsetBuffer::<i32>::from_lengths(vec![2, 3]),
316            None,
317        );
318
319        let res = keep_greater_than_two(list).unwrap();
320        let actual = res.as_list::<i32>();
321
322        // [1,5] -> [5], [2,4,3] -> [4,3]
323        let expected = create_i32_list(
324            vec![5, 4, 3],
325            OffsetBuffer::<i32>::from_lengths(vec![1, 2]),
326            None,
327        );
328
329        assert_eq!(actual, &expected);
330    }
331
332    #[test]
333    fn filter_on_sliced_list_should_not_evaluate_on_unreachable_values() {
334        // First sublist [0] is sliced away; sliced array covers sublists [1..3]
335        let list = create_i32_list(
336            vec![
337                0, // unreachable after slice — if evaluated, it would appear in output
338                1, 5, 2, 4, 3, 7,
339            ],
340            OffsetBuffer::<i32>::from_lengths(vec![1, 3, 3]),
341            None,
342        )
343        .slice(1, 2);
344
345        let res = keep_greater_than_two(list).unwrap();
346        let actual = res.as_list::<i32>();
347
348        // [1,5,2] -> [5], [4,3,7] -> [4,3,7]
349        let expected = create_i32_list(
350            vec![5, 4, 3, 7],
351            OffsetBuffer::<i32>::from_lengths(vec![1, 3]),
352            None,
353        );
354
355        assert_eq!(actual, &expected);
356    }
357
358    #[test]
359    fn filter_should_not_be_evaluated_on_values_underlying_null() {
360        // The null sublist (index 1) contains values that would pass the predicate
361        // if evaluated. We verify they do NOT appear in the output.
362        let list = create_i32_list(
363            vec![1, 5, 99, 100, 3, 7],
364            OffsetBuffer::<i32>::from_lengths(vec![2, 2, 2]),
365            Some(NullBuffer::from(vec![true, false, true])),
366        );
367
368        let res = keep_greater_than_two(list).unwrap();
369        let actual = res.as_list::<i32>();
370
371        // sublist 0: [1,5] -> [5]
372        // sublist 1: null  -> null (empty range, null bit)
373        // sublist 2: [3,7] -> [3,7]
374        let expected = create_i32_list(
375            vec![5, 3, 7],
376            OffsetBuffer::<i32>::from_lengths(vec![1, 0, 2]),
377            Some(NullBuffer::from(vec![true, false, true])),
378        );
379
380        assert_eq!(actual.data_type(), expected.data_type());
381        assert_eq!(actual, &expected);
382    }
383
384    #[test]
385    fn filter_all_filtered_out() {
386        let list =
387            create_i32_list(vec![1, 2], OffsetBuffer::<i32>::from_lengths(vec![2]), None);
388
389        let res = keep_greater_than_two(list).unwrap();
390        let actual = res.as_list::<i32>();
391
392        let expected = create_i32_list(
393            vec![0i32; 0],
394            OffsetBuffer::<i32>::from_lengths(vec![0]),
395            None,
396        );
397
398        assert_eq!(actual, &expected);
399    }
400
401    #[test]
402    fn filter_nothing_filtered_reuses_values() {
403        let list = create_i32_list(
404            vec![3, 4, 5],
405            OffsetBuffer::<i32>::from_lengths(vec![3]),
406            None,
407        );
408        // all elements > 2, so nothing is filtered — values buffer should be reused
409        let res = keep_greater_than_two(list.clone()).unwrap();
410        assert_eq!(res.as_list::<i32>(), &list);
411    }
412
413    #[test]
414    fn scalar_true_predicate_returns_original_list() {
415        let list = create_i32_list(
416            vec![1, 2, 3],
417            OffsetBuffer::<i32>::from_lengths(vec![3]),
418            None,
419        );
420        // x -> true: every element kept, should return list unchanged
421        let res = eval_hof_on_i32_list(
422            array_filter_higher_order_function(),
423            list.clone(),
424            lit(true),
425        )
426        .unwrap();
427        assert_eq!(res.as_list::<i32>(), &list);
428    }
429
430    #[test]
431    fn scalar_false_predicate_returns_empty_sublists() {
432        let list = create_i32_list(
433            vec![1, 2, 3, 4],
434            OffsetBuffer::<i32>::from_lengths(vec![2, 2]),
435            None,
436        );
437        // x -> false: every sublist emptied
438        let res =
439            eval_hof_on_i32_list(array_filter_higher_order_function(), list, lit(false))
440                .unwrap();
441        let actual = res.as_list::<i32>();
442        let expected = create_i32_list(
443            vec![0i32; 0],
444            OffsetBuffer::<i32>::from_lengths(vec![0, 0]),
445            None,
446        );
447        assert_eq!(actual, &expected);
448    }
449
450    #[test]
451    fn filter_large_list_parity() {
452        let list = create_i32_large_list(
453            vec![1, 2, 3, 4, 5],
454            OffsetBuffer::<i64>::from_lengths(vec![5]),
455            None,
456        );
457        let res = keep_greater_than_two(list).unwrap();
458        let actual = res.as_list::<i64>();
459        let expected = create_i32_large_list(
460            vec![3, 4, 5],
461            OffsetBuffer::<i64>::from_lengths(vec![3]),
462            None,
463        );
464        assert_eq!(actual, &expected);
465    }
466
467    #[test]
468    fn filter_captured_outer_column() {
469        let list = create_i32_list(
470            vec![1, 50, 4, 50, 7, 50],
471            OffsetBuffer::<i32>::from_lengths(vec![2, 2, 2]),
472            None,
473        );
474        let number = Int32Array::from(vec![10, 40, 60]);
475        let res = eval_hof_on_i32_list_with_outer(
476            array_filter_higher_order_function(),
477            list,
478            number,
479            v().gt(col("number")),
480        )
481        .unwrap();
482        let actual = res.as_list::<i32>();
483        let expected = create_i32_list(
484            vec![50, 50],
485            OffsetBuffer::<i32>::from_lengths(vec![1, 1, 0]),
486            None,
487        );
488        assert_eq!(actual, &expected);
489    }
490}