diff --git a/datafusion/functions-nested/src/sort.rs b/datafusion/functions-nested/src/sort.rs index f4f9148f760bf..b70304ffcf65c 100644 --- a/datafusion/functions-nested/src/sort.rs +++ b/datafusion/functions-nested/src/sort.rs @@ -18,11 +18,11 @@ //! [`ScalarUDFImpl`] definitions for array_sort function. use crate::utils::make_scalar_function; -use arrow::array::BooleanBufferBuilder; use arrow::array::{ Array, ArrayRef, ArrowPrimitiveType, GenericListArray, OffsetSizeTrait, - PrimitiveArray, UInt32Array, UInt64Array, new_empty_array, new_null_array, + PrimitiveArray, UInt32Array, UInt64Array, new_empty_array, }; +use arrow::array::{BooleanBufferBuilder, StringArray}; use arrow::buffer::{NullBuffer, OffsetBuffer}; use arrow::datatypes::{ArrowNativeTypeOp, DataType, FieldRef}; use arrow::row::{RowConverter, SortField}; @@ -162,22 +162,14 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result { return Ok(Arc::clone(&args[0])); } - if args[1..].iter().any(|array| array.is_null(0)) { - return Ok(new_null_array(args[0].data_type(), args[0].len())); - } + let sort_order = if args.len() > 1 { + Some(as_string_array(&args[1])?) + } else { + None + }; - let sort_options = if args.len() >= 2 { - let order = as_string_array(&args[1])?.value(0); - let descending = order_desc(order)?; - let nulls_first = if args.len() >= 3 { - order_nulls_first(as_string_array(&args[2])?.value(0))? - } else { - true - }; - Some(SortOptions { - descending, - nulls_first, - }) + let null_order = if args.len() > 2 { + Some(as_string_array(&args[2])?) } else { None }; @@ -190,11 +182,11 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result { } DataType::List(field) => { let array = as_list_array(&args[0])?; - array_sort_generic(array, Arc::clone(field), sort_options) + array_sort_generic(array, Arc::clone(field), sort_order, null_order) } DataType::LargeList(field) => { let array = as_large_list_array(&args[0])?; - array_sort_generic(array, Arc::clone(field), sort_options) + array_sort_generic(array, Arc::clone(field), sort_order, null_order) } // Signature should prevent this arm ever occurring _ => exec_err!("array_sort expects list for first argument"), @@ -204,14 +196,15 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result { fn array_sort_generic( list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, + null_order: Option<&StringArray>, ) -> Result { let values = list_array.values(); if values.data_type().is_primitive() { - array_sort_primitive(list_array, field, sort_options) + array_sort_primitive(list_array, field, sort_order, null_order) } else { - array_sort_non_primitive(list_array, field, sort_options) + array_sort_non_primitive(list_array, field, sort_order, null_order) } } @@ -220,11 +213,12 @@ fn array_sort_generic( fn array_sort_primitive( list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, + null_order: Option<&StringArray>, ) -> Result { let values = list_array.values().as_ref(); downcast_primitive_array! { - values => sort_primitive_list(values, list_array, field, sort_options), + values => sort_primitive_list(values, list_array, field, sort_order, null_order), _ => exec_err!("array_sort: unsupported primitive type") } } @@ -233,15 +227,16 @@ fn sort_primitive_list( prim_values: &PrimitiveArray, list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, + null_order: Option<&StringArray>, ) -> Result where T::Native: ArrowNativeTypeOp, { if prim_values.null_count() > 0 { - sort_list_with_nulls(prim_values, list_array, field, sort_options) + sort_list_with_nulls(prim_values, list_array, field, sort_order, null_order) } else { - sort_list_no_nulls(prim_values, list_array, field, sort_options) + sort_list_no_nulls(prim_values, list_array, field, sort_order) } } @@ -251,7 +246,7 @@ fn sort_list_no_nulls( prim_values: &PrimitiveArray, list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, ) -> Result where T::Native: ArrowNativeTypeOp, @@ -261,8 +256,6 @@ where let values_start = offsets[0].as_usize(); let values_end = offsets[row_count].as_usize(); - let descending = sort_options.is_some_and(|o| o.descending); - // Copy all values into a mutable buffer let mut values: Vec = prim_values.values()[values_start..values_end].to_vec(); @@ -274,6 +267,14 @@ where let start = window[0].as_usize() - values_start; let end = window[1].as_usize() - values_start; let slice = &mut values[start..end]; + let descending = if let Some(sort_order) = sort_order { + if sort_order.is_null(row_index) { + continue; + } + order_desc(sort_order.value(row_index))? + } else { + false + }; if descending { slice.sort_unstable_by(|a, b| b.compare(*a)); } else { @@ -300,7 +301,8 @@ fn sort_list_with_nulls( prim_values: &PrimitiveArray, list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, + null_order: Option<&StringArray>, ) -> Result where T::Native: ArrowNativeTypeOp, @@ -310,9 +312,7 @@ where let values_start = offsets[0].as_usize(); let values_end = offsets[row_count].as_usize(); let total_values = values_end - values_start; - - let descending = sort_options.is_some_and(|o| o.descending); - let nulls_first = sort_options.is_none_or(|o| o.nulls_first); + let mut list_validity = BooleanBufferBuilder::new(row_count); let mut out_values: Vec = vec![T::Native::default(); total_values]; let mut validity = BooleanBufferBuilder::new(total_values); @@ -332,12 +332,24 @@ where if list_array.is_null(row_index) || row_len == 0 { validity.append_n(row_len, false); + list_validity.append(false); continue; } let null_count = src_nulls.slice(start, row_len).null_count(); let valid_count = row_len - null_count; + let nulls_first = if let Some(null_order) = null_order { + if null_order.is_null(row_index) { + list_validity.append(false); + validity.append_n(row_len, false); + continue; + } + order_nulls_first(null_order.value(row_index))? + } else { + true + }; + // Compact valid values directly into the target region of the output // buffer: after nulls (if nulls_first) or at the start (if nulls_last). let valid_offset = if nulls_first { null_count } else { 0 }; @@ -351,6 +363,16 @@ where let valid_slice = &mut out_values [out_start + valid_offset..out_start + valid_offset + valid_count]; + let descending = if let Some(sort_order) = sort_order { + if sort_order.is_null(row_index) { + validity.append_n(row_len, false); + list_validity.append(false); + continue; + } + order_desc(sort_order.value(row_index))? + } else { + false + }; if descending { valid_slice.sort_unstable_by(|a, b| b.compare(*a)); } else { @@ -365,6 +387,8 @@ where validity.append_n(valid_count, true); validity.append_n(null_count, false); } + + list_validity.append(true); } let new_offsets = rebase_offsets(offsets); @@ -379,7 +403,7 @@ where field, new_offsets, sorted_values, - list_array.nulls().cloned(), + Some(NullBuffer::from(list_validity.finish())), )?)) } @@ -390,7 +414,8 @@ where fn array_sort_non_primitive( list_array: &GenericListArray, field: FieldRef, - sort_options: Option, + sort_order: Option<&StringArray>, + null_order: Option<&StringArray>, ) -> Result { let row_count = list_array.len(); let values = list_array.values(); @@ -398,12 +423,47 @@ fn array_sort_non_primitive( let values_start = offsets[0].as_usize(); let total_values = offsets[row_count].as_usize() - values_start; - let converter = RowConverter::new(vec![SortField::new_with_options( + let desc_first_converter = RowConverter::new(vec![SortField::new_with_options( values.data_type().clone(), - sort_options.unwrap_or_default(), + SortOptions { + descending: true, + nulls_first: true, + }, )])?; + + let desc_last_converter = RowConverter::new(vec![SortField::new_with_options( + values.data_type().clone(), + SortOptions { + descending: true, + nulls_first: false, + }, + )])?; + + let asc_first_converter = RowConverter::new(vec![SortField::new_with_options( + values.data_type().clone(), + SortOptions { + descending: false, + nulls_first: true, + }, + )])?; + + let asc_last_converter = RowConverter::new(vec![SortField::new_with_options( + values.data_type().clone(), + SortOptions { + descending: false, + nulls_first: false, + }, + )])?; + let values_sliced = values.slice(values_start, total_values); - let rows = converter.convert_columns(&[Arc::clone(&values_sliced)])?; + let desc_first_rows = + desc_first_converter.convert_columns(&[Arc::clone(&values_sliced)])?; + let desc_last_rows = + desc_last_converter.convert_columns(&[Arc::clone(&values_sliced)])?; + let asc_first_rows = + asc_first_converter.convert_columns(&[Arc::clone(&values_sliced)])?; + let asc_last_rows = + asc_last_converter.convert_columns(&[Arc::clone(&values_sliced)])?; let mut indices: Vec = Vec::with_capacity(total_values); let mut new_offsets = Vec::with_capacity(row_count + 1); @@ -420,6 +480,26 @@ fn array_sort_non_primitive( continue; } + let descending = if let Some(sort_order) = sort_order { + if sort_order.is_null(row_index) { + new_offsets.push(new_offsets[row_index]); + continue; + } + order_desc(sort_order.value(row_index))? + } else { + false + }; + + let nulls_first = if let Some(null_order) = null_order { + if null_order.is_null(row_index) { + new_offsets.push(new_offsets[row_index]); + continue; + } + order_nulls_first(null_order.value(row_index))? + } else { + true + }; + let len = (end - start).as_usize(); let local_start = start.as_usize() - values_start; @@ -428,7 +508,29 @@ fn array_sort_non_primitive( } else { sort_scratch.clear(); sort_scratch.extend(local_start..local_start + len); - sort_scratch.sort_unstable_by(|&a, &b| rows.row(a).cmp(&rows.row(b))); + + if descending { + if nulls_first { + sort_scratch.sort_unstable_by(|&a, &b| { + desc_first_rows.row(a).cmp(&desc_first_rows.row(b)) + }); + } else { + sort_scratch.sort_unstable_by(|&a, &b| { + desc_last_rows.row(a).cmp(&desc_last_rows.row(b)) + }); + } + } else { + if nulls_first { + sort_scratch.sort_unstable_by(|&a, &b| { + asc_first_rows.row(a).cmp(&asc_first_rows.row(b)) + }); + } else { + sort_scratch.sort_unstable_by(|&a, &b| { + asc_last_rows.row(a).cmp(&asc_last_rows.row(b)) + }); + } + } + indices.extend(sort_scratch.iter().map(|&i| OffsetSize::usize_as(i))); } diff --git a/datafusion/sqllogictest/test_files/array/array_sort.slt b/datafusion/sqllogictest/test_files/array/array_sort.slt index 343aa5a82ef0c..caa160a856d06 100644 --- a/datafusion/sqllogictest/test_files/array/array_sort.slt +++ b/datafusion/sqllogictest/test_files/array/array_sort.slt @@ -264,5 +264,61 @@ select array_sort(arrow_cast([1, 3, null, 5, NULL, -5], 'LargeListView(Int64)')) ---- [NULL, NULL, -5, 1, 3, 5] +# Ensure that sort options (asc/desc, null placement) can vary per-row. + +query ? +select array_sort(arr, order) +from values + (make_array(5, 2, 3), 'asc'), + (make_array(5, 2, 3), 'desc'), + (make_array(10, 20, 30), 'desc') as t(arr, order); +---- +[2, 3, 5] +[5, 3, 2] +[30, 20, 10] + +query ? +select array_sort(arr, order) +from values + (make_array(5, 2, 3), 'asc'), + (make_array(5, NULL, 3), 'desc'), + (make_array(10, 20, 30), 'desc') as t(arr, order); +---- +[2, 3, 5] +[NULL, 5, 3] +[30, 20, 10] + +query ? +select array_sort(arr, order) +from values + (make_array('b', 'a'), 'asc'), + (make_array('y', 'z', 'x'), 'desc'), + (make_array('m', 'n', NULL), 'desc') as t(arr, order); +---- +[a, b] +[z, y, x] +[NULL, n, m] + +query ? +select array_sort(arr, order, null_order) +from values + (make_array(5, 2, NULL, 3, NULL), 'asc', 'NULLS FIRST'), + (make_array(NULL, 5, NULL, 3, null), 'desc', 'NULLS LAST'), + (make_array(10, NULL, null, 20, 30, NULL), 'desc', 'NULLS LAST') as t(arr, order, null_order); +---- +[NULL, NULL, 2, 3, 5] +[5, 3, NULL, NULL, NULL] +[30, 20, 10, NULL, NULL, NULL] + +query ? +select array_sort(arr, order, null_order) +from values + (make_array('a', 'c', NULL, 'b', NULL), 'asc', 'NULLS FIRST'), + (make_array(NULL, 'x', NULL, 'y', null), 'desc', 'NULLS LAST'), + (make_array('j', NULL, null, 'l', 'k', NULL), 'desc', 'NULLS LAST') as t(arr, order, null_order); +---- +[NULL, NULL, a, b, c] +[y, x, NULL, NULL, NULL] +[l, k, j, NULL, NULL, NULL] include ./cleanup.slt.part