1use 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 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 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
204fn 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
232fn 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 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 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 let list = create_i32_list(
336 vec![
337 0, 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 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 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 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 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 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 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}