Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 110 additions & 1 deletion datafusion/functions-nested/src/min_max.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
use crate::utils::make_scalar_function;
use arrow::array::{
Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType, AsArray, GenericListArray,
OffsetSizeTrait, PrimitiveBuilder, downcast_primitive,
OffsetSizeTrait, PrimitiveBuilder, downcast_primitive, new_empty_array,
};
use arrow::buffer::NullBuffer;
use arrow::datatypes::DataType;
Expand Down Expand Up @@ -209,6 +209,12 @@ fn array_min_max_helper<O: OffsetSizeTrait>(
return result;
}

// `ScalarValue::iter_to_array` below cannot build an array from an empty
// iterator, so return an empty array of the element type for zero rows.
if array.is_empty() {
return Ok(new_empty_array(&array.value_type()));
}

// Fallback: per-row ScalarValue path for non-primitive types
let agg_fn = if is_min { min_batch } else { max_batch };
let null_value = ScalarValue::try_from(array.value_type())?;
Expand Down Expand Up @@ -311,3 +317,106 @@ fn scalar_min_max<N: ArrowNativeTypeOp>(
}
best
}

#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{GenericListBuilder, ListArray, StringArray, StringBuilder};
use arrow::datatypes::Int64Type;

fn utf8_list<O: OffsetSizeTrait>(
rows: Vec<Option<Vec<Option<&str>>>>,
) -> GenericListArray<O> {
let mut builder = GenericListBuilder::<O, _>::new(StringBuilder::new());
for row in rows {
match row {
Some(values) => {
for v in values {
builder.values().append_option(v);
}
builder.append(true);
}
None => builder.append(false),
}
}
builder.finish()
}

#[test]
fn zero_rows_non_primitive_returns_empty_array() -> Result<()> {
for is_min in [true, false] {
let out = array_min_max_helper(&utf8_list::<i32>(vec![]), is_min)?;
assert_eq!(out.len(), 0);
assert_eq!(out.data_type(), &DataType::Utf8);

let out = array_min_max_helper(&utf8_list::<i64>(vec![]), is_min)?;
assert_eq!(out.len(), 0);
assert_eq!(out.data_type(), &DataType::Utf8);
}
Ok(())
}

#[test]
fn zero_rows_primitive_returns_empty_array() -> Result<()> {
let empty = ListArray::from_iter_primitive::<Int64Type, _, _>(Vec::<
Option<Vec<Option<i64>>>,
>::new());
for is_min in [true, false] {
let out = array_min_max_helper(&empty, is_min)?;
assert_eq!(out.len(), 0);
assert_eq!(out.data_type(), &DataType::Int64);
}
Ok(())
}

#[test]
fn non_primitive_rows() -> Result<()> {
let list = utf8_list::<i32>(vec![
Some(vec![Some("prod"), Some("api")]),
Some(vec![]),
None,
Some(vec![Some("web"), None, Some("db")]),
]);
let min = array_min_max_helper(&list, true)?;
let max = array_min_max_helper(&list, false)?;
assert_eq!(
min.as_string::<i32>(),
&StringArray::from(vec![Some("api"), None, None, Some("db")])
);
assert_eq!(
max.as_string::<i32>(),
&StringArray::from(vec![Some("prod"), None, None, Some("web")])
);
Ok(())
}

#[test]
fn invoke_on_zero_row_batch() -> Result<()> {
// What ProjectionExec / TopK do when an upstream operator emits a
// zero-row batch.
let list: ArrayRef = Arc::new(utf8_list::<i32>(vec![]));
for udf in [array_min_udf(), array_max_udf()] {
let out = udf.invoke_with_args(ScalarFunctionArgs {
args: vec![ColumnarValue::Array(Arc::clone(&list))],
arg_fields: vec![Arc::new(arrow::datatypes::Field::new(
"a",
list.data_type().clone(),
true,
))],
number_rows: 0,
return_field: Arc::new(arrow::datatypes::Field::new(
"r",
DataType::Utf8,
true,
)),
config_options: Arc::new(Default::default()),
})?;
let ColumnarValue::Array(out) = out else {
panic!("expected an array")
};
assert_eq!(out.len(), 0);
assert_eq!(out.data_type(), &DataType::Utf8);
}
Ok(())
}
}