use std::sync::Arc;
use arrow::datatypes::{DataType, Field, FieldRef, Fields};
use arrow::array::{
Array, ArrayRef, BooleanArray, Float64Array, GenericListArray, NullBufferBuilder,
OffsetSizeTrait, Scalar,
};
use arrow::buffer::{NullBuffer, OffsetBuffer};
use datafusion_common::cast::{
as_fixed_size_list_array, as_float64_array, as_generic_list_array,
as_large_list_array, as_large_list_view_array, as_list_array, as_list_view_array,
};
use datafusion_common::{Result, ScalarValue, exec_err, internal_err, plan_err};
use datafusion_expr::ColumnarValue;
use itertools::Itertools as _;
pub(crate) fn list_type_with_element(
array_type: &DataType,
element_nullable: bool,
) -> DataType {
match array_type {
DataType::List(field) => {
DataType::List(widen_nullability(field, element_nullable))
}
DataType::LargeList(field) => {
DataType::LargeList(widen_nullability(field, element_nullable))
}
other => other.clone(),
}
}
fn widen_nullability(field: &FieldRef, nullable: bool) -> FieldRef {
if nullable && !field.is_nullable() {
Arc::new(field.as_ref().clone().with_nullable(true))
} else {
Arc::clone(field)
}
}
pub(crate) fn list_inner_field(context: &str, data_type: &DataType) -> Result<FieldRef> {
match data_type {
DataType::List(field) | DataType::LargeList(field) => Ok(Arc::clone(field)),
other => internal_err!("{context} got unexpected data type: {other}"),
}
}
pub(crate) fn check_datatypes(name: &str, args: &[&ArrayRef]) -> Result<()> {
let data_type = args[0].data_type();
if !args.iter().all(|arg| {
arg.data_type().equals_datatype(data_type)
|| arg.data_type().equals_datatype(&DataType::Null)
}) {
let types = args.iter().map(|arg| arg.data_type()).collect::<Vec<_>>();
return plan_err!(
"{name} received incompatible types: {}",
types.iter().join(", ")
);
}
Ok(())
}
pub(crate) fn make_scalar_function<F>(
inner: F,
) -> impl Fn(&[ColumnarValue]) -> Result<ColumnarValue>
where
F: Fn(&[ArrayRef]) -> Result<ArrayRef>,
{
move |args: &[ColumnarValue]| {
let len = args
.iter()
.fold(Option::<usize>::None, |acc, arg| match arg {
ColumnarValue::Scalar(_) => acc,
ColumnarValue::Array(a) => Some(a.len()),
});
let is_scalar = len.is_none();
let args = ColumnarValue::values_to_arrays(args)?;
let result = (inner)(&args);
if is_scalar {
let result = result.and_then(|arr| ScalarValue::try_from_array(&arr, 0));
result.map(ColumnarValue::Scalar)
} else {
result.map(ColumnarValue::Array)
}
}
}
pub(crate) fn align_array_dimensions<O: OffsetSizeTrait>(
args: Vec<ArrayRef>,
) -> Result<Vec<ArrayRef>> {
let args_ndim = args
.iter()
.map(|arg| datafusion_common::utils::list_ndims(arg.data_type()))
.collect::<Vec<_>>();
let max_ndim = args_ndim.iter().max().unwrap_or(&0);
let aligned_args: Result<Vec<ArrayRef>> = args
.into_iter()
.zip(args_ndim.iter())
.map(|(array, ndim)| {
if ndim < max_ndim {
let mut aligned_array = Arc::clone(&array);
for _ in 0..(max_ndim - ndim) {
let data_type = aligned_array.data_type().to_owned();
let array_lengths = vec![1; aligned_array.len()];
let offsets = OffsetBuffer::<O>::from_lengths(array_lengths);
aligned_array = Arc::new(GenericListArray::<O>::try_new(
Arc::new(Field::new_list_field(data_type, true)),
offsets,
aligned_array,
None,
)?)
}
Ok(aligned_array)
} else {
Ok(Arc::clone(&array))
}
})
.collect();
aligned_args
}
pub(crate) fn compare_element_to_list(
list_array_row: &dyn Array,
element_array: &dyn Array,
row_index: usize,
eq: bool,
) -> Result<BooleanArray> {
if list_array_row.data_type() != element_array.data_type() {
return exec_err!(
"compare_element_to_list received incompatible types: '{:?}' and '{:?}'.",
list_array_row.data_type(),
element_array.data_type()
);
}
let element_array_row = element_array.slice(row_index, 1);
let res = match element_array_row.data_type() {
DataType::List(_) => {
let element_array_row_inner = as_list_array(&element_array_row)?.value(0);
let list_array_row_inner = as_list_array(list_array_row)?;
list_array_row_inner
.iter()
.map(|row| {
row.map(|row| {
if eq {
row.eq(&element_array_row_inner)
} else {
row.ne(&element_array_row_inner)
}
})
})
.collect::<BooleanArray>()
}
DataType::LargeList(_) => {
let element_array_row_inner =
as_large_list_array(&element_array_row)?.value(0);
let list_array_row_inner = as_large_list_array(list_array_row)?;
list_array_row_inner
.iter()
.map(|row| {
row.map(|row| {
if eq {
row.eq(&element_array_row_inner)
} else {
row.ne(&element_array_row_inner)
}
})
})
.collect::<BooleanArray>()
}
_ => {
let element_arr = Scalar::new(element_array_row);
if eq {
arrow_ord::cmp::not_distinct(&list_array_row, &element_arr)?
} else {
arrow_ord::cmp::distinct(&list_array_row, &element_arr)?
}
}
};
Ok(res)
}
pub(crate) fn compute_array_dims(
arr: Option<ArrayRef>,
) -> Result<Option<Vec<Option<u64>>>> {
let mut value = match arr {
Some(arr) => arr,
None => return Ok(None),
};
if value.is_empty() {
return Ok(None);
}
let mut res = vec![Some(value.len() as u64)];
loop {
match value.data_type() {
DataType::List(_) => {
value = as_list_array(&value)?.value(0);
res.push(Some(value.len() as u64));
}
DataType::LargeList(_) => {
value = as_large_list_array(&value)?.value(0);
res.push(Some(value.len() as u64));
}
DataType::ListView(_) => {
value = as_list_view_array(&value)?.value(0);
res.push(Some(value.len() as u64));
}
DataType::LargeListView(_) => {
value = as_large_list_view_array(&value)?.value(0);
res.push(Some(value.len() as u64));
}
DataType::FixedSizeList(..) => {
value = as_fixed_size_list_array(&value)?.value(0);
res.push(Some(value.len() as u64));
}
_ => return Ok(Some(res)),
}
}
}
pub(crate) fn get_map_entry_field(data_type: &DataType) -> Result<&Fields> {
match data_type {
DataType::Map(field, _) => {
let field_data_type = field.data_type();
match field_data_type {
DataType::Struct(fields) => Ok(fields),
_ => {
internal_err!("Expected a Struct type, got {}", field_data_type)
}
}
}
_ => internal_err!("Expected a Map type, got {data_type}"),
}
}
pub(crate) fn coerce_array_math_arg_types(
name: &str,
arg_types: &[DataType],
) -> Result<Vec<DataType>> {
use DataType::{FixedSizeList, LargeList, List, Null};
use datafusion_common::utils::{ListCoercion, coerced_type_with_base_type_only};
let coercion = Some(&ListCoercion::FixedSizedListToList);
for arg_type in arg_types {
if !matches!(arg_type, Null | List(_) | LargeList(_) | FixedSizeList(..)) {
return plan_err!("{name} does not support type {arg_type}");
}
}
let any_large_list = arg_types.iter().any(|t| matches!(t, LargeList(_)));
let coerced = arg_types
.iter()
.map(|arg_type| {
if matches!(arg_type, Null) {
let field = Arc::new(Field::new_list_field(DataType::Float64, true));
return if any_large_list {
LargeList(field)
} else {
List(field)
};
}
let coerced =
coerced_type_with_base_type_only(arg_type, &DataType::Float64, coercion);
match coerced {
List(field) if any_large_list => LargeList(field),
other => other,
}
})
.collect();
Ok(coerced)
}
pub(crate) fn array_math_binary_op<O, F>(
op_name: &str,
lhs: &ArrayRef,
rhs: &ArrayRef,
op: F,
) -> Result<ArrayRef>
where
O: OffsetSizeTrait,
F: Fn(f64, f64) -> f64,
{
let lhs = as_generic_list_array::<O>(lhs)?;
let rhs = as_generic_list_array::<O>(rhs)?;
let lhs_values = as_float64_array(lhs.values())?;
let rhs_values = as_float64_array(rhs.values())?;
let lhs_offsets = lhs.value_offsets();
let rhs_offsets = rhs.value_offsets();
let row_nulls = NullBuffer::union(lhs.nulls(), rhs.nulls());
let mut out_values: Vec<f64> = Vec::with_capacity(lhs_values.len());
let mut out_inner_nulls = NullBufferBuilder::new(lhs_values.len());
let mut out_offsets = Vec::<O>::with_capacity(lhs.len() + 1);
out_offsets.push(O::zero());
for row in 0..lhs.len() {
if row_nulls.as_ref().is_some_and(|nb| nb.is_null(row)) {
out_offsets.push(out_offsets[row]);
continue;
}
let start1 = lhs_offsets[row].as_usize();
let len1 = lhs.value_length(row).as_usize();
let start2 = rhs_offsets[row].as_usize();
let len2 = rhs.value_length(row).as_usize();
if len1 != len2 {
return exec_err!(
"{op_name} requires both list inputs to have the same length per row, got {len1} and {len2} at row {row}"
);
}
let l_slice = lhs_values.slice(start1, len1);
let r_slice = rhs_values.slice(start2, len2);
let l_vals = l_slice.values();
let r_vals = r_slice.values();
for i in 0..len1 {
out_values.push(op(l_vals[i], r_vals[i]));
}
match NullBuffer::union(l_slice.nulls(), r_slice.nulls()) {
Some(nb) => out_inner_nulls.append_buffer(&nb),
None => out_inner_nulls.append_n_non_nulls(len1),
}
out_offsets.push(out_offsets[row] + O::usize_as(len1));
}
let values_array = Arc::new(Float64Array::new(
out_values.into(),
out_inner_nulls.finish(),
));
let field = Arc::new(Field::new_list_field(DataType::Float64, true));
Ok(Arc::new(GenericListArray::<O>::try_new(
field,
OffsetBuffer::new(out_offsets.into()),
values_array,
row_nulls,
)?))
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::ListArray;
use arrow::datatypes::Int64Type;
use datafusion_common::utils::SingleRowListArrayBuilder;
#[test]
fn test_align_array_dimensions() {
let array1d_1: ArrayRef =
Arc::new(ListArray::from_iter_primitive::<Int64Type, _, _>(vec![
Some(vec![Some(1), Some(2), Some(3)]),
Some(vec![Some(4), Some(5)]),
]));
let array1d_2: ArrayRef =
Arc::new(ListArray::from_iter_primitive::<Int64Type, _, _>(vec![
Some(vec![Some(6), Some(7), Some(8)]),
]));
let array2d_1: ArrayRef = Arc::new(
SingleRowListArrayBuilder::new(Arc::clone(&array1d_1)).build_list_array(),
);
let array2d_2 = Arc::new(
SingleRowListArrayBuilder::new(Arc::clone(&array1d_2)).build_list_array(),
);
let res = align_array_dimensions::<i32>(vec![
array1d_1.to_owned(),
array2d_2.to_owned(),
])
.unwrap();
let expected = as_list_array(&array2d_1).unwrap();
let expected_dim = datafusion_common::utils::list_ndims(array2d_1.data_type());
assert_ne!(as_list_array(&res[0]).unwrap(), expected);
assert_eq!(
datafusion_common::utils::list_ndims(res[0].data_type()),
expected_dim
);
let array3d_1: ArrayRef =
Arc::new(SingleRowListArrayBuilder::new(array2d_1).build_list_array());
let array3d_2: ArrayRef =
Arc::new(SingleRowListArrayBuilder::new(array2d_2).build_list_array());
let res = align_array_dimensions::<i32>(vec![array1d_1, array3d_2]).unwrap();
let expected = as_list_array(&array3d_1).unwrap();
let expected_dim = datafusion_common::utils::list_ndims(array3d_1.data_type());
assert_ne!(as_list_array(&res[0]).unwrap(), expected);
assert_eq!(
datafusion_common::utils::list_ndims(res[0].data_type()),
expected_dim
);
}
}