|
19 | 19 | use crate::utils::make_scalar_function; |
20 | 20 | use arrow::array::{ |
21 | 21 | Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType, AsArray, GenericListArray, |
22 | | - OffsetSizeTrait, PrimitiveBuilder, downcast_primitive, |
| 22 | + OffsetSizeTrait, PrimitiveBuilder, downcast_primitive, new_empty_array, |
23 | 23 | }; |
24 | 24 | use arrow::buffer::NullBuffer; |
25 | 25 | use arrow::datatypes::DataType; |
@@ -209,6 +209,12 @@ fn array_min_max_helper<O: OffsetSizeTrait>( |
209 | 209 | return result; |
210 | 210 | } |
211 | 211 |
|
| 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 | + |
212 | 218 | // Fallback: per-row ScalarValue path for non-primitive types |
213 | 219 | let agg_fn = if is_min { min_batch } else { max_batch }; |
214 | 220 | let null_value = ScalarValue::try_from(array.value_type())?; |
@@ -311,3 +317,106 @@ fn scalar_min_max<N: ArrowNativeTypeOp>( |
311 | 317 | } |
312 | 318 | best |
313 | 319 | } |
| 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