#![forbid(unsafe_code)]
use std::sync::Arc;
use arrow::array::{
Array, ArrayRef, AsArray, BooleanArray, BooleanBuilder, Int64Builder, new_null_array,
};
use arrow::buffer::NullBuffer;
use arrow::compute::take;
use arrow::compute::take_arrays;
use arrow::datatypes::{ArrowNativeType, DataType, Field, FieldRef};
use datafusion::common::utils::{
adjust_offsets_for_slice, list_values, list_values_row_number, take_function_args,
};
use datafusion::error::DataFusionError;
use datafusion::logical_expr::{
ColumnarValue, HigherOrderFunctionArgs, HigherOrderReturnFieldArgs, HigherOrderSignature,
HigherOrderUDF, HigherOrderUDFImpl, LambdaParametersProgress, ValueOrLambda, Volatility,
};
use datafusion::prelude::SessionContext;
use datafusion_functions_nested::array_any_match::ArrayAnyMatch;
use datafusion_functions_nested::array_filter::ArrayFilter;
use datafusion_functions_nested::array_transform::ArrayTransform;
type DFResult<T> = Result<T, DataFusionError>;
pub fn register_higher_order_spark_functions(ctx: &SessionContext) -> DFResult<()> {
ctx.register_higher_order_function(Arc::new(
HigherOrderUDF::new_from_impl(ArrayTransform::new()).with_aliases(["transform"]),
));
ctx.register_higher_order_function(Arc::new(
HigherOrderUDF::new_from_impl(ArrayFilter::new()).with_aliases(["filter"]),
));
ctx.register_higher_order_function(Arc::new(
HigherOrderUDF::new_from_impl(ArrayAnyMatch::new()).with_aliases(["exists"]),
));
ctx.register_higher_order_function(Arc::new(HigherOrderUDF::new_from_impl(
ArrayAllMatch::new(),
)));
ctx.register_higher_order_function(Arc::new(HigherOrderUDF::new_from_impl(ArrayReduce::new())));
Ok(())
}
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct ArrayReduce {
signature: HigherOrderSignature,
aliases: Vec<String>,
}
impl Default for ArrayReduce {
fn default() -> Self {
Self::new()
}
}
impl ArrayReduce {
pub fn new() -> Self {
Self {
signature: HigherOrderSignature::exact(
vec![
ValueOrLambda::Value(()),
ValueOrLambda::Value(()),
ValueOrLambda::Lambda(()),
],
Volatility::Immutable,
),
aliases: vec![
String::from("aggregate"),
String::from("reduce"),
String::from("array_aggregate"),
],
}
}
}
impl HigherOrderUDFImpl for ArrayReduce {
fn name(&self) -> &str {
"array_reduce"
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn signature(&self) -> &HigherOrderSignature {
&self.signature
}
fn coerce_value_types(&self, arg_types: &[DataType]) -> DFResult<Vec<DataType>> {
let [list, init] = take_function_args(self.name(), arg_types)?;
let coerced_list = match list {
DataType::List(_) | DataType::LargeList(_) => list.clone(),
DataType::ListView(field) | DataType::FixedSizeList(field, _) => {
DataType::List(Arc::clone(field))
}
DataType::LargeListView(field) => DataType::LargeList(Arc::clone(field)),
other => {
return Err(DataFusionError::Plan(format!(
"{} expected a list as first argument, got {other}",
self.name()
)));
}
};
Ok(vec![coerced_list, init.clone()])
}
fn lambda_parameters(
&self,
_step: usize,
fields: &[ValueOrLambda<FieldRef, Option<FieldRef>>],
) -> DFResult<LambdaParametersProgress> {
let [list, init, _merge] = take_function_args(self.name(), fields)?;
let (ValueOrLambda::Value(list), ValueOrLambda::Value(init)) = (list, init) else {
return Err(DataFusionError::Plan(format!(
"{} expects two value arguments before the lambda",
self.name()
)));
};
let element = match list.data_type() {
DataType::List(field) | DataType::LargeList(field) => Arc::clone(field),
other => {
return Err(DataFusionError::Plan(format!("expected list, got {other}")));
}
};
Ok(LambdaParametersProgress::Complete(vec![vec![
Arc::clone(init),
element,
]]))
}
fn return_field_from_args(&self, args: HigherOrderReturnFieldArgs) -> DFResult<Arc<Field>> {
let [_list, init, merge] = take_function_args(self.name(), args.arg_fields)?;
let data_type = match merge {
ValueOrLambda::Lambda(field) => field.data_type().clone(),
ValueOrLambda::Value(_) => match init {
ValueOrLambda::Value(field) => field.data_type().clone(),
ValueOrLambda::Lambda(_) => {
return Err(DataFusionError::Plan(format!(
"{} expects a start value as the second argument",
self.name()
)));
}
},
};
Ok(Arc::new(Field::new("", data_type, true)))
}
fn invoke_with_args(&self, args: HigherOrderFunctionArgs) -> DFResult<ColumnarValue> {
let num_rows = args.number_rows;
let [list, init, merge] = take_function_args(self.name(), &args.args)?;
let (ValueOrLambda::Value(list), ValueOrLambda::Value(init), ValueOrLambda::Lambda(merge)) =
(list, init, merge)
else {
return Err(DataFusionError::Execution(format!(
"{} expects (array, start, lambda)",
self.name()
)));
};
let list_array = list.to_array(num_rows)?;
let mut acc = init.to_array(num_rows)?;
let (offsets, values): (Vec<i64>, ArrayRef) = match list_array.data_type() {
DataType::List(_) => {
let list = list_array.as_list::<i32>();
(
list.offsets()
.iter()
.map(|offset| i64::from(*offset))
.collect(),
Arc::clone(list.values()),
)
}
DataType::LargeList(_) => {
let list = list_array.as_list::<i64>();
(
list.offsets().iter().copied().collect(),
Arc::clone(list.values()),
)
}
other => {
return Err(DataFusionError::Execution(format!(
"expected list, got {other}"
)));
}
};
let lengths: Vec<i64> = offsets
.windows(2)
.map(|window| match window {
[start, end] => *end - *start,
_ => 0,
})
.collect();
let max_len = lengths.iter().copied().max().unwrap_or(0);
for k in 0..max_len {
let mut index_builder = Int64Builder::with_capacity(num_rows);
let mut has_kth = BooleanBuilder::with_capacity(num_rows);
for (offset, len) in offsets.iter().zip(lengths.iter()) {
if k < *len {
index_builder.append_value(*offset + k);
has_kth.append_value(true);
} else {
index_builder.append_null();
has_kth.append_value(false);
}
}
let indices = index_builder.finish();
let has_kth = has_kth.finish();
let element = take(values.as_ref(), &indices, None)?;
let accumulator = Arc::clone(&acc);
let acc_fn: &dyn Fn() -> DFResult<ArrayRef> = &|| Ok(Arc::clone(&accumulator));
let element_fn: &dyn Fn() -> DFResult<ArrayRef> = &|| Ok(Arc::clone(&element));
let merged = merge
.evaluate(&[acc_fn, element_fn], |arrays| Ok(arrays.to_vec()))?
.into_array(num_rows)?;
let merged_ref: &dyn Array = merged.as_ref();
let acc_ref: &dyn Array = acc.as_ref();
acc = arrow::compute::kernels::zip::zip(&has_kth, &merged_ref, &acc_ref)?;
}
if let Some(nulls) = list_array.nulls() {
let valid = BooleanArray::new(nulls.inner().clone(), None);
let null_array = new_null_array(acc.data_type(), num_rows);
let acc_ref: &dyn Array = acc.as_ref();
let null_ref: &dyn Array = null_array.as_ref();
acc = arrow::compute::kernels::zip::zip(&valid, &acc_ref, &null_ref)?;
}
Ok(ColumnarValue::Array(acc))
}
}
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct ArrayAllMatch {
signature: HigherOrderSignature,
aliases: Vec<String>,
}
impl Default for ArrayAllMatch {
fn default() -> Self {
Self::new()
}
}
impl ArrayAllMatch {
pub fn new() -> Self {
Self {
signature: HigherOrderSignature::exact(
vec![ValueOrLambda::Value(()), ValueOrLambda::Lambda(())],
Volatility::Immutable,
),
aliases: vec![String::from("forall"), String::from("array_forall")],
}
}
}
fn all_match_for_range(predicate: &BooleanArray, start: usize, end: usize) -> Option<bool> {
let any_false = (start..end).any(|j| predicate.is_valid(j) && !predicate.value(j));
if any_false {
return Some(false);
}
let any_null = (start..end).any(|j| predicate.is_null(j));
if any_null { None } else { Some(true) }
}
impl HigherOrderUDFImpl for ArrayAllMatch {
fn name(&self) -> &str {
"array_all_match"
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn signature(&self) -> &HigherOrderSignature {
&self.signature
}
fn coerce_value_types(&self, arg_types: &[DataType]) -> DFResult<Vec<DataType>> {
let [list] = arg_types else {
return Err(DataFusionError::Plan(format!(
"{} requires 1 value argument, got {}",
self.name(),
arg_types.len()
)));
};
let coerced = match list {
DataType::List(_) | DataType::LargeList(_) => list.clone(),
DataType::ListView(field) | DataType::FixedSizeList(field, _) => {
DataType::List(Arc::clone(field))
}
DataType::LargeListView(field) => DataType::LargeList(Arc::clone(field)),
other => {
return Err(DataFusionError::Plan(format!(
"{} expected a list as first argument, got {other}",
self.name()
)));
}
};
Ok(vec![coerced])
}
fn lambda_parameters(
&self,
_step: usize,
fields: &[ValueOrLambda<FieldRef, Option<FieldRef>>],
) -> DFResult<LambdaParametersProgress> {
let [list, _] = take_function_args(self.name(), fields)?;
let ValueOrLambda::Value(list) = list else {
return Err(DataFusionError::Plan(format!(
"{} expects a value as first argument",
self.name()
)));
};
let field = match list.data_type() {
DataType::List(field) | DataType::LargeList(field) => field,
other => {
return Err(DataFusionError::Plan(format!("expected list, got {other}")));
}
};
Ok(LambdaParametersProgress::Complete(vec![vec![Arc::clone(
field,
)]]))
}
fn return_field_from_args(&self, args: HigherOrderReturnFieldArgs) -> DFResult<Arc<Field>> {
let [ValueOrLambda::Value(list), _] = take_function_args(self.name(), args.arg_fields)?
else {
return Err(DataFusionError::Plan(format!(
"{} expects a value as first argument",
self.name()
)));
};
Ok(Arc::new(Field::new(
"",
DataType::Boolean,
list.is_nullable(),
)))
}
fn invoke_with_args(&self, args: HigherOrderFunctionArgs) -> DFResult<ColumnarValue> {
let [ValueOrLambda::Value(list), ValueOrLambda::Lambda(lambda)] =
take_function_args(self.name(), &args.args)?
else {
return Err(DataFusionError::Execution(format!(
"{} expects a value followed by a lambda",
self.name()
)));
};
let list_array = list.to_array(args.number_rows)?;
if list_array.null_count() == list_array.len() {
return Ok(ColumnarValue::Array(new_null_array(
args.return_type(),
list_array.len(),
)));
}
let list_values = list_values(&list_array)?;
let values_param = || Ok(Arc::clone(&list_values));
let predicate_results = lambda
.evaluate(&[&values_param], |arrays| {
let indices = list_values_row_number(&list_array)?;
Ok(take_arrays(arrays, &indices, None)?)
})?
.into_array(list_values.len())?;
let predicate_bool = predicate_results
.as_any()
.downcast_ref::<BooleanArray>()
.ok_or_else(|| {
DataFusionError::Execution(format!(
"{} predicate must return a boolean array",
self.name()
))
})?;
let mut values = BooleanBuilder::with_capacity(list_array.len());
macro_rules! process_list {
($list_typed:expr) => {{
let offsets = adjust_offsets_for_slice($list_typed);
let offsets: &[_] = &offsets;
for pair in offsets.windows(2) {
let [start, end] = pair else { continue };
values.append_option(all_match_for_range(
predicate_bool,
start.as_usize(),
end.as_usize(),
));
}
}};
}
match list_array.data_type() {
DataType::List(_) => process_list!(list_array.as_list::<i32>()),
DataType::LargeList(_) => process_list!(list_array.as_list::<i64>()),
other => {
return Err(DataFusionError::Execution(format!(
"expected list, got {other}"
)));
}
}
let (boolean_buffer, predicate_nulls) = values.finish().into_parts();
let nulls = NullBuffer::union(list_array.nulls(), predicate_nulls.as_ref());
Ok(ColumnarValue::Array(Arc::new(BooleanArray::new(
boolean_buffer,
nulls,
))))
}
}
#[cfg(test)]
mod tests {
use arrow::array::{Array, BooleanArray, Int64Array};
async fn run(sql: &str) -> Vec<arrow::array::RecordBatch> {
crate::SqlEngine::new()
.sql(sql)
.await
.expect("plan")
.collect()
.await
.expect("collect")
}
#[tokio::test]
async fn spark_transform_alias_doubles_elements() {
let b = run("SELECT transform([1, 2, 3], x -> x * 2) AS r").await;
let list = b[0]
.column(0)
.as_any()
.downcast_ref::<arrow::array::ListArray>();
let list = list.expect("list");
let vals = list.value(0);
let vals = vals.as_any().downcast_ref::<Int64Array>().expect("i64");
assert_eq!(vals.values(), &[2, 4, 6]);
}
#[tokio::test]
async fn spark_filter_alias_keeps_matching() {
let b = run("SELECT filter([1, 2, 3, 4], x -> x % 2 = 0) AS r").await;
let list = b[0]
.column(0)
.as_any()
.downcast_ref::<arrow::array::ListArray>()
.expect("list");
let vals = list.value(0);
let vals = vals.as_any().downcast_ref::<Int64Array>().expect("i64");
assert_eq!(vals.values(), &[2, 4]);
}
#[tokio::test]
async fn spark_exists_and_forall() {
let b = run("SELECT any_match([1, 2, 3], x -> x > 2) AS any_gt2, \
forall([2, 4, 6], x -> x % 2 = 0) AS all_even, \
forall([2, 3, 6], x -> x % 2 = 0) AS not_all_even, \
forall(filter([1], x -> x > 100), x -> x > 0) AS empty_all")
.await;
let row = &b[0];
let col = |i: usize| {
row.column(i)
.as_any()
.downcast_ref::<BooleanArray>()
.expect("bool")
.value(0)
};
assert!(col(0), "exists any > 2");
assert!(col(1), "forall even");
assert!(!col(2), "not all even");
assert!(col(3), "forall over empty array is true");
}
#[tokio::test]
async fn forall_null_semantics() {
let b = run(
"SELECT forall([2, 4], x -> CASE WHEN x = 4 THEN NULL ELSE x % 2 = 0 END) AS poisoned, \
forall([3, 4], x -> CASE WHEN x = 4 THEN NULL ELSE x % 2 = 0 END) AS false_wins",
)
.await;
let row = &b[0];
let poisoned = row
.column(0)
.as_any()
.downcast_ref::<BooleanArray>()
.unwrap();
let false_wins = row
.column(1)
.as_any()
.downcast_ref::<BooleanArray>()
.unwrap();
assert!(poisoned.is_null(0), "NULL result with no false ⇒ NULL");
assert!(
!false_wins.is_null(0) && !false_wins.value(0),
"false dominates NULL"
);
}
#[tokio::test]
async fn spark_aggregate_left_folds_sum() {
let b = run(
"SELECT aggregate([1, 2, 3, 4], 0, (acc, x) -> acc + x) AS s, \
reduce([1, 2, 3], 10, (acc, x) -> acc + x) AS s10",
)
.await;
let row = &b[0];
let col = |i: usize| {
row.column(i)
.as_any()
.downcast_ref::<Int64Array>()
.expect("i64")
.value(0)
};
assert_eq!(col(0), 10, "0 + 1 + 2 + 3 + 4");
assert_eq!(col(1), 16, "10 + 1 + 2 + 3");
}
#[tokio::test]
async fn aggregate_multi_row_varying_lengths_and_null() {
let b = run(
"SELECT aggregate(arr, 0, (acc, x) -> acc + x) AS s FROM (VALUES \
(1, [1, 2, 3]), \
(2, [10]), \
(3, CAST(NULL AS INT[])), \
(4, [5, 5, 5, 5]) \
) AS t(id, arr) ORDER BY id",
)
.await;
let s = b[0]
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.expect("i64");
assert_eq!(s.value(0), 6, "1+2+3");
assert_eq!(s.value(1), 10, "single element");
assert!(s.is_null(2), "NULL array ⇒ NULL result");
assert_eq!(s.value(3), 20, "5*4");
}
#[tokio::test]
async fn all_match_range_helper_direct() {
use super::all_match_for_range;
let p = BooleanArray::from(vec![Some(true), Some(true), Some(false), None]);
assert_eq!(all_match_for_range(&p, 0, 2), Some(true));
assert_eq!(all_match_for_range(&p, 0, 3), Some(false)); assert_eq!(all_match_for_range(&p, 0, 4), Some(false)); assert_eq!(all_match_for_range(&p, 3, 4), None); assert_eq!(all_match_for_range(&p, 1, 1), Some(true)); }
}