Skip to content

Commit 1dd1fbd

Browse files
izveigoralamb
andauthored
feat: array_contains (#6618)
* feat: array_contains * feat: regen.sh * docs: array_contains * fix: merge * Update docs/source/user-guide/sql/scalar_functions.md --------- Co-authored-by: Andrew Lamb <andrew@nerdnetworks.org>
1 parent 4f2933f commit 1dd1fbd

12 files changed

Lines changed: 229 additions & 25 deletions

File tree

datafusion/core/tests/sqllogictests/test_files/array.slt

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -357,3 +357,39 @@ query ?
357357
select make_array(x, y) from foo2;
358358
----
359359
[1.0, 1]
360+
361+
# array_contains scalar function #1
362+
query BBB rowsort
363+
select array_contains(make_array(1, 2, 3), make_array(1, 1, 2, 3)), array_contains([1, 2, 3], [1, 1, 2]), array_contains([1, 2, 3], [2, 1, 3, 1]);
364+
----
365+
true true true
366+
367+
# array_contains scalar function #2
368+
query BB rowsort
369+
select array_contains([[1, 2], [3, 4]], [[1, 2], [3, 4], [1, 3]]), array_contains([[[1], [2]], [[3], [4]]], [1, 2, 2, 3, 4]);
370+
----
371+
true true
372+
373+
# array_contains scalar function #3
374+
query BBB rowsort
375+
select array_contains(make_array(1, 2, 3), make_array(1, 2, 3, 4)), array_contains([1, 2, 3], [1, 1, 4]), array_contains([1, 2, 3], [2, 1, 3, 4]);
376+
----
377+
false false false
378+
379+
# array_contains scalar function #4
380+
query BB rowsort
381+
select array_contains([[1, 2], [3, 4]], [[1, 2], [3, 4], [1, 5]]), array_contains([[[1], [2]], [[3], [4]]], [1, 2, 2, 3, 5]);
382+
----
383+
false false
384+
385+
# array_contains scalar function #5
386+
query BB rowsort
387+
select array_contains([true, true, false, true, false], [true, false, false]), array_contains([true, false, true], [true, true]);
388+
----
389+
true true
390+
391+
# array_contains scalar function #6
392+
query BB rowsort
393+
select array_contains(make_array(true, true, true), make_array(false, false)), array_contains([false, false, false], [true, true]);
394+
----
395+
false false

datafusion/expr/src/built_in_function.rs

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -113,6 +113,8 @@ pub enum BuiltinScalarFunction {
113113
ArrayAppend,
114114
/// array_concat
115115
ArrayConcat,
116+
/// array_contains
117+
ArrayContains,
116118
/// array_dims
117119
ArrayDims,
118120
/// array_fill
@@ -319,6 +321,7 @@ impl BuiltinScalarFunction {
319321
BuiltinScalarFunction::Trunc => Volatility::Immutable,
320322
BuiltinScalarFunction::ArrayAppend => Volatility::Immutable,
321323
BuiltinScalarFunction::ArrayConcat => Volatility::Immutable,
324+
BuiltinScalarFunction::ArrayContains => Volatility::Immutable,
322325
BuiltinScalarFunction::ArrayDims => Volatility::Immutable,
323326
BuiltinScalarFunction::ArrayFill => Volatility::Immutable,
324327
BuiltinScalarFunction::ArrayLength => Volatility::Immutable,
@@ -460,6 +463,7 @@ impl BuiltinScalarFunction {
460463
"The {self} function can only accept fixed size list as the args."
461464
))),
462465
},
466+
BuiltinScalarFunction::ArrayContains => Ok(Boolean),
463467
BuiltinScalarFunction::ArrayDims => Ok(UInt8),
464468
BuiltinScalarFunction::ArrayFill => Ok(List(Arc::new(Field::new(
465469
"item",
@@ -741,6 +745,7 @@ impl BuiltinScalarFunction {
741745
BuiltinScalarFunction::ArrayConcat => {
742746
Signature::variadic_any(self.volatility())
743747
}
748+
BuiltinScalarFunction::ArrayContains => Signature::any(2, self.volatility()),
744749
BuiltinScalarFunction::ArrayDims => Signature::any(1, self.volatility()),
745750
BuiltinScalarFunction::ArrayFill => Signature::any(2, self.volatility()),
746751
BuiltinScalarFunction::ArrayLength => {
@@ -1166,6 +1171,7 @@ fn aliases(func: &BuiltinScalarFunction) -> &'static [&'static str] {
11661171
// array functions
11671172
BuiltinScalarFunction::ArrayAppend => &["array_append"],
11681173
BuiltinScalarFunction::ArrayConcat => &["array_concat"],
1174+
BuiltinScalarFunction::ArrayContains => &["array_contains"],
11691175
BuiltinScalarFunction::ArrayDims => &["array_dims"],
11701176
BuiltinScalarFunction::ArrayFill => &["array_fill"],
11711177
BuiltinScalarFunction::ArrayLength => &["array_length"],

datafusion/expr/src/expr_fn.rs

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -530,6 +530,13 @@ scalar_expr!(
530530
"appends an element to the end of an array."
531531
);
532532
nary_scalar_expr!(ArrayConcat, array_concat, "concatenates arrays.");
533+
scalar_expr!(
534+
ArrayContains,
535+
array_contains,
536+
first_array second_array,
537+
"returns true, if each element of the second array appe
538+
aring in the first array, otherwise false."
539+
);
533540
scalar_expr!(
534541
ArrayDims,
535542
array_dims,

datafusion/physical-expr/src/array_expressions.rs

Lines changed: 124 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@ use datafusion_common::cast::as_list_array;
2626
use datafusion_common::ScalarValue;
2727
use datafusion_common::{DataFusionError, Result};
2828
use datafusion_expr::ColumnarValue;
29+
use itertools::Itertools;
2930
use std::sync::Arc;
3031

3132
macro_rules! downcast_vec {
@@ -1070,6 +1071,70 @@ pub fn array_ndims(args: &[ColumnarValue]) -> Result<ColumnarValue> {
10701071
]))))
10711072
}
10721073

1074+
macro_rules! contains {
1075+
($FIRST_ARRAY:expr, $SECOND_ARRAY:expr, $ARRAY_TYPE:ident) => {{
1076+
let first_array = downcast_arg!($FIRST_ARRAY, $ARRAY_TYPE);
1077+
let second_array = downcast_arg!($SECOND_ARRAY, $ARRAY_TYPE);
1078+
let mut res = true;
1079+
for x in second_array.values().iter().dedup() {
1080+
if !first_array.values().contains(x) {
1081+
res = false;
1082+
}
1083+
}
1084+
1085+
res
1086+
}};
1087+
}
1088+
1089+
/// Array_contains SQL function
1090+
pub fn array_contains(args: &[ArrayRef]) -> Result<ArrayRef> {
1091+
fn concat_inner_lists(arg: ArrayRef) -> Result<ArrayRef> {
1092+
match arg.data_type() {
1093+
DataType::List(field) => match field.data_type() {
1094+
DataType::List(..) => {
1095+
concat_inner_lists(array_concat(&[as_list_array(&arg)?
1096+
.values()
1097+
.clone()])?)
1098+
}
1099+
_ => Ok(as_list_array(&arg)?.values().clone()),
1100+
},
1101+
data_type => Err(DataFusionError::NotImplemented(format!(
1102+
"Array is not type '{data_type:?}'."
1103+
))),
1104+
}
1105+
}
1106+
1107+
let concat_first_array = concat_inner_lists(args[0].clone())?.clone();
1108+
let concat_second_array = concat_inner_lists(args[1].clone())?.clone();
1109+
1110+
let res = match (concat_first_array.data_type(), concat_second_array.data_type()) {
1111+
(DataType::Utf8, DataType::Utf8) => contains!(concat_first_array, concat_second_array, StringArray),
1112+
(DataType::LargeUtf8, DataType::LargeUtf8) => contains!(concat_first_array, concat_second_array, LargeStringArray),
1113+
(DataType::Boolean, DataType::Boolean) => {
1114+
let first_array = downcast_arg!(concat_first_array, BooleanArray);
1115+
let second_array = downcast_arg!(concat_second_array, BooleanArray);
1116+
compute::bool_or(first_array) == compute::bool_or(second_array)
1117+
}
1118+
(DataType::Float32, DataType::Float32) => contains!(concat_first_array, concat_second_array, Float32Array),
1119+
(DataType::Float64, DataType::Float64) => contains!(concat_first_array, concat_second_array, Float64Array),
1120+
(DataType::Int8, DataType::Int8) => contains!(concat_first_array, concat_second_array, Int8Array),
1121+
(DataType::Int16, DataType::Int16) => contains!(concat_first_array, concat_second_array, Int16Array),
1122+
(DataType::Int32, DataType::Int32) => contains!(concat_first_array, concat_second_array, Int32Array),
1123+
(DataType::Int64, DataType::Int64) => contains!(concat_first_array, concat_second_array, Int64Array),
1124+
(DataType::UInt8, DataType::UInt8) => contains!(concat_first_array, concat_second_array, UInt8Array),
1125+
(DataType::UInt16, DataType::UInt16) => contains!(concat_first_array, concat_second_array, UInt16Array),
1126+
(DataType::UInt32, DataType::UInt32) => contains!(concat_first_array, concat_second_array, UInt32Array),
1127+
(DataType::UInt64, DataType::UInt64) => contains!(concat_first_array, concat_second_array, UInt64Array),
1128+
(first_array_data_type, second_array_data_type) => {
1129+
return Err(DataFusionError::NotImplemented(format!(
1130+
"Array_contains is not implemented for types '{first_array_data_type:?}' and '{second_array_data_type:?}'."
1131+
)))
1132+
}
1133+
};
1134+
1135+
Ok(Arc::new(BooleanArray::from(vec![res])))
1136+
}
1137+
10731138
#[cfg(test)]
10741139
mod tests {
10751140
use super::*;
@@ -1588,7 +1653,7 @@ mod tests {
15881653

15891654
#[test]
15901655
fn test_array_ndims() {
1591-
// array_ndims([1, 2]) = 1
1656+
// array_ndims([1, 2, 3, 4]) = 1
15921657
let list_array = return_array();
15931658

15941659
let array = array_ndims(&[list_array])
@@ -1602,7 +1667,7 @@ mod tests {
16021667

16031668
#[test]
16041669
fn test_nested_array_ndims() {
1605-
// array_ndims([[1, 2], [3, 4]]) = 2
1670+
// array_ndims([[1, 2, 3, 4], [5, 6, 7, 8]]) = 2
16061671
let list_array = return_nested_array();
16071672

16081673
let array = array_ndims(&[list_array])
@@ -1614,6 +1679,63 @@ mod tests {
16141679
assert_eq!(result, &UInt8Array::from(vec![2]));
16151680
}
16161681

1682+
#[test]
1683+
fn test_array_contains() {
1684+
// array_contains([1, 2, 3, 4], array_append([1, 2, 3, 4], 3)) = t
1685+
let first_array = return_array().into_array(1);
1686+
let second_array = array_append(&[
1687+
first_array.clone(),
1688+
Arc::new(Int64Array::from(vec![Some(3)])),
1689+
])
1690+
.expect("failed to initialize function array_contains");
1691+
1692+
let arr = array_contains(&[first_array.clone(), second_array])
1693+
.expect("failed to initialize function array_contains");
1694+
let result = as_boolean_array(&arr);
1695+
1696+
assert_eq!(result, &BooleanArray::from(vec![true]));
1697+
1698+
// array_contains([1, 2, 3, 4], array_append([1, 2, 3, 4], 5)) = f
1699+
let second_array = array_append(&[
1700+
first_array.clone(),
1701+
Arc::new(Int64Array::from(vec![Some(5)])),
1702+
])
1703+
.expect("failed to initialize function array_contains");
1704+
1705+
let arr = array_contains(&[first_array.clone(), second_array])
1706+
.expect("failed to initialize function array_contains");
1707+
let result = as_boolean_array(&arr);
1708+
1709+
assert_eq!(result, &BooleanArray::from(vec![false]));
1710+
}
1711+
1712+
#[test]
1713+
fn test_nested_array_contains() {
1714+
// array_contains([[1, 2, 3, 4], [5, 6, 7, 8]], array_append([1, 2, 3, 4], 3)) = t
1715+
let first_array = return_nested_array().into_array(1);
1716+
let array = return_array().into_array(1);
1717+
let second_array =
1718+
array_append(&[array.clone(), Arc::new(Int64Array::from(vec![Some(3)]))])
1719+
.expect("failed to initialize function array_contains");
1720+
1721+
let arr = array_contains(&[first_array.clone(), second_array])
1722+
.expect("failed to initialize function array_contains");
1723+
let result = as_boolean_array(&arr);
1724+
1725+
assert_eq!(result, &BooleanArray::from(vec![true]));
1726+
1727+
// array_contains([[1, 2, 3, 4], [5, 6, 7, 8]], array_append([1, 2, 3, 4], 9)) = f
1728+
let second_array =
1729+
array_append(&[array.clone(), Arc::new(Int64Array::from(vec![Some(9)]))])
1730+
.expect("failed to initialize function array_contains");
1731+
1732+
let arr = array_contains(&[first_array.clone(), second_array])
1733+
.expect("failed to initialize function array_contains");
1734+
let result = as_boolean_array(&arr);
1735+
1736+
assert_eq!(result, &BooleanArray::from(vec![false]));
1737+
}
1738+
16171739
fn return_array() -> ColumnarValue {
16181740
let args = [
16191741
ColumnarValue::Scalar(ScalarValue::Int64(Some(1))),

datafusion/physical-expr/src/functions.rs

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -391,6 +391,9 @@ pub fn create_physical_fun(
391391
BuiltinScalarFunction::ArrayConcat => {
392392
Arc::new(|args| make_scalar_function(array_expressions::array_concat)(args))
393393
}
394+
BuiltinScalarFunction::ArrayContains => {
395+
Arc::new(|args| make_scalar_function(array_expressions::array_contains)(args))
396+
}
394397
BuiltinScalarFunction::ArrayDims => Arc::new(array_expressions::array_dims),
395398
BuiltinScalarFunction::ArrayFill => Arc::new(array_expressions::array_fill),
396399
BuiltinScalarFunction::ArrayLength => Arc::new(array_expressions::array_length),

datafusion/proto/proto/datafusion.proto

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -563,6 +563,7 @@ enum ScalarFunction {
563563
ArrayToString = 97;
564564
Cardinality = 98;
565565
TrimArray = 99;
566+
ArrayContains = 100;
566567
}
567568

568569
message ScalarFunctionNode {

datafusion/proto/src/generated/pbjson.rs

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

datafusion/proto/src/generated/prost.rs

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

datafusion/proto/src/logical_plan/from_proto.rs

Lines changed: 11 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -36,12 +36,12 @@ use datafusion_common::{
3636
};
3737
use datafusion_expr::expr::Placeholder;
3838
use datafusion_expr::{
39-
abs, acos, acosh, array, array_append, array_concat, array_dims, array_fill,
40-
array_length, array_ndims, array_position, array_positions, array_prepend,
41-
array_remove, array_replace, array_to_string, ascii, asin, asinh, atan, atan2, atanh,
42-
bit_length, btrim, cardinality, cbrt, ceil, character_length, chr, coalesce,
43-
concat_expr, concat_ws_expr, cos, cosh, date_bin, date_part, date_trunc, degrees,
44-
digest, exp,
39+
abs, acos, acosh, array, array_append, array_concat, array_contains, array_dims,
40+
array_fill, array_length, array_ndims, array_position, array_positions,
41+
array_prepend, array_remove, array_replace, array_to_string, ascii, asin, asinh,
42+
atan, atan2, atanh, bit_length, btrim, cardinality, cbrt, ceil, character_length,
43+
chr, coalesce, concat_expr, concat_ws_expr, cos, cosh, date_bin, date_part,
44+
date_trunc, degrees, digest, exp,
4545
expr::{self, InList, Sort, WindowFunction},
4646
factorial, floor, from_unixtime, gcd, lcm, left, ln, log, log10, log2,
4747
logical_plan::{PlanType, StringifiedPlan},
@@ -450,6 +450,7 @@ impl From<&protobuf::ScalarFunction> for BuiltinScalarFunction {
450450
ScalarFunction::ToTimestamp => Self::ToTimestamp,
451451
ScalarFunction::ArrayAppend => Self::ArrayAppend,
452452
ScalarFunction::ArrayConcat => Self::ArrayConcat,
453+
ScalarFunction::ArrayContains => Self::ArrayContains,
453454
ScalarFunction::ArrayDims => Self::ArrayDims,
454455
ScalarFunction::ArrayFill => Self::ArrayFill,
455456
ScalarFunction::ArrayLength => Self::ArrayLength,
@@ -1192,6 +1193,10 @@ pub fn parse_expr(
11921193
.map(|expr| parse_expr(expr, registry))
11931194
.collect::<Result<Vec<_>, _>>()?,
11941195
)),
1196+
ScalarFunction::ArrayContains => Ok(array_contains(
1197+
parse_expr(&args[0], registry)?,
1198+
parse_expr(&args[1], registry)?,
1199+
)),
11951200
ScalarFunction::ArrayFill => Ok(array_fill(
11961201
parse_expr(&args[0], registry)?,
11971202
parse_expr(&args[1], registry)?,

datafusion/proto/src/logical_plan/to_proto.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1344,6 +1344,7 @@ impl TryFrom<&BuiltinScalarFunction> for protobuf::ScalarFunction {
13441344
BuiltinScalarFunction::ToTimestamp => Self::ToTimestamp,
13451345
BuiltinScalarFunction::ArrayAppend => Self::ArrayAppend,
13461346
BuiltinScalarFunction::ArrayConcat => Self::ArrayConcat,
1347+
BuiltinScalarFunction::ArrayContains => Self::ArrayContains,
13471348
BuiltinScalarFunction::ArrayDims => Self::ArrayDims,
13481349
BuiltinScalarFunction::ArrayFill => Self::ArrayFill,
13491350
BuiltinScalarFunction::ArrayLength => Self::ArrayLength,

0 commit comments

Comments
 (0)