Skip to content

Commit 3b29882

Browse files
committed
fix: array_min/array_max on zero rows for non-primitive element types
1 parent d55153a commit 3b29882

1 file changed

Lines changed: 110 additions & 1 deletion

File tree

‎datafusion/functions-nested/src/min_max.rs‎

Lines changed: 110 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
use crate::utils::make_scalar_function;
2020
use arrow::array::{
2121
Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType, AsArray, GenericListArray,
22-
OffsetSizeTrait, PrimitiveBuilder, downcast_primitive,
22+
OffsetSizeTrait, PrimitiveBuilder, downcast_primitive, new_empty_array,
2323
};
2424
use arrow::buffer::NullBuffer;
2525
use arrow::datatypes::DataType;
@@ -209,6 +209,12 @@ fn array_min_max_helper<O: OffsetSizeTrait>(
209209
return result;
210210
}
211211

212+
// `ScalarValue::iter_to_array` below cannot build an array from an empty
213+
// iterator, so return an empty array of the element type for zero rows.
214+
if array.is_empty() {
215+
return Ok(new_empty_array(&array.value_type()));
216+
}
217+
212218
// Fallback: per-row ScalarValue path for non-primitive types
213219
let agg_fn = if is_min { min_batch } else { max_batch };
214220
let null_value = ScalarValue::try_from(array.value_type())?;
@@ -311,3 +317,106 @@ fn scalar_min_max<N: ArrowNativeTypeOp>(
311317
}
312318
best
313319
}
320+
321+
#[cfg(test)]
322+
mod tests {
323+
use super::*;
324+
use arrow::array::{GenericListBuilder, ListArray, StringArray, StringBuilder};
325+
use arrow::datatypes::Int64Type;
326+
327+
fn utf8_list<O: OffsetSizeTrait>(
328+
rows: Vec<Option<Vec<Option<&str>>>>,
329+
) -> GenericListArray<O> {
330+
let mut builder = GenericListBuilder::<O, _>::new(StringBuilder::new());
331+
for row in rows {
332+
match row {
333+
Some(values) => {
334+
for v in values {
335+
builder.values().append_option(v);
336+
}
337+
builder.append(true);
338+
}
339+
None => builder.append(false),
340+
}
341+
}
342+
builder.finish()
343+
}
344+
345+
#[test]
346+
fn zero_rows_non_primitive_returns_empty_array() -> Result<()> {
347+
for is_min in [true, false] {
348+
let out = array_min_max_helper(&utf8_list::<i32>(vec![]), is_min)?;
349+
assert_eq!(out.len(), 0);
350+
assert_eq!(out.data_type(), &DataType::Utf8);
351+
352+
let out = array_min_max_helper(&utf8_list::<i64>(vec![]), is_min)?;
353+
assert_eq!(out.len(), 0);
354+
assert_eq!(out.data_type(), &DataType::Utf8);
355+
}
356+
Ok(())
357+
}
358+
359+
#[test]
360+
fn zero_rows_primitive_returns_empty_array() -> Result<()> {
361+
let empty = ListArray::from_iter_primitive::<Int64Type, _, _>(Vec::<
362+
Option<Vec<Option<i64>>>,
363+
>::new());
364+
for is_min in [true, false] {
365+
let out = array_min_max_helper(&empty, is_min)?;
366+
assert_eq!(out.len(), 0);
367+
assert_eq!(out.data_type(), &DataType::Int64);
368+
}
369+
Ok(())
370+
}
371+
372+
#[test]
373+
fn non_primitive_rows() -> Result<()> {
374+
let list = utf8_list::<i32>(vec![
375+
Some(vec![Some("prod"), Some("api")]),
376+
Some(vec![]),
377+
None,
378+
Some(vec![Some("web"), None, Some("db")]),
379+
]);
380+
let min = array_min_max_helper(&list, true)?;
381+
let max = array_min_max_helper(&list, false)?;
382+
assert_eq!(
383+
min.as_string::<i32>(),
384+
&StringArray::from(vec![Some("api"), None, None, Some("db")])
385+
);
386+
assert_eq!(
387+
max.as_string::<i32>(),
388+
&StringArray::from(vec![Some("prod"), None, None, Some("web")])
389+
);
390+
Ok(())
391+
}
392+
393+
#[test]
394+
fn invoke_on_zero_row_batch() -> Result<()> {
395+
// What ProjectionExec / TopK do when an upstream operator emits a
396+
// zero-row batch.
397+
let list: ArrayRef = Arc::new(utf8_list::<i32>(vec![]));
398+
for udf in [array_min_udf(), array_max_udf()] {
399+
let out = udf.invoke_with_args(ScalarFunctionArgs {
400+
args: vec![ColumnarValue::Array(Arc::clone(&list))],
401+
arg_fields: vec![Arc::new(arrow::datatypes::Field::new(
402+
"a",
403+
list.data_type().clone(),
404+
true,
405+
))],
406+
number_rows: 0,
407+
return_field: Arc::new(arrow::datatypes::Field::new(
408+
"r",
409+
DataType::Utf8,
410+
true,
411+
)),
412+
config_options: Arc::new(Default::default()),
413+
})?;
414+
let ColumnarValue::Array(out) = out else {
415+
panic!("expected an array")
416+
};
417+
assert_eq!(out.len(), 0);
418+
assert_eq!(out.data_type(), &DataType::Utf8);
419+
}
420+
Ok(())
421+
}
422+
}

0 commit comments

Comments
 (0)