Skip to content
Draft
Show file tree
Hide file tree
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
182 changes: 142 additions & 40 deletions datafusion/functions-nested/src/sort.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -162,22 +162,14 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result<ArrayRef> {
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
};
Expand All @@ -190,11 +182,11 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result<ArrayRef> {
}
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"),
Expand All @@ -204,14 +196,15 @@ fn array_sort_inner(args: &[ArrayRef]) -> Result<ArrayRef> {
fn array_sort_generic<OffsetSize: OffsetSizeTrait>(
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
null_order: Option<&StringArray>,
) -> Result<ArrayRef> {
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)
}
}

Expand All @@ -220,11 +213,12 @@ fn array_sort_generic<OffsetSize: OffsetSizeTrait>(
fn array_sort_primitive<OffsetSize: OffsetSizeTrait>(
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
null_order: Option<&StringArray>,
) -> Result<ArrayRef> {
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")
}
}
Expand All @@ -233,15 +227,16 @@ fn sort_primitive_list<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
prim_values: &PrimitiveArray<T>,
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
null_order: Option<&StringArray>,
) -> Result<ArrayRef>
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)
}
}

Expand All @@ -251,7 +246,7 @@ fn sort_list_no_nulls<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
prim_values: &PrimitiveArray<T>,
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
) -> Result<ArrayRef>
where
T::Native: ArrowNativeTypeOp,
Expand All @@ -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<T::Native> =
prim_values.values()[values_start..values_end].to_vec();
Expand All @@ -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 {
Expand All @@ -300,7 +301,8 @@ fn sort_list_with_nulls<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
prim_values: &PrimitiveArray<T>,
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
null_order: Option<&StringArray>,
) -> Result<ArrayRef>
where
T::Native: ArrowNativeTypeOp,
Expand All @@ -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<T::Native> = vec![T::Native::default(); total_values];
let mut validity = BooleanBufferBuilder::new(total_values);
Expand All @@ -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 };
Expand All @@ -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 {
Expand All @@ -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);
Expand All @@ -379,7 +403,7 @@ where
field,
new_offsets,
sorted_values,
list_array.nulls().cloned(),
Some(NullBuffer::from(list_validity.finish())),
)?))
}

Expand All @@ -390,20 +414,56 @@ where
fn array_sort_non_primitive<OffsetSize: OffsetSizeTrait>(
list_array: &GenericListArray<OffsetSize>,
field: FieldRef,
sort_options: Option<SortOptions>,
sort_order: Option<&StringArray>,
null_order: Option<&StringArray>,
) -> Result<ArrayRef> {
let row_count = list_array.len();
let values = list_array.values();
let offsets = list_array.offsets();
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<OffsetSize> = Vec::with_capacity(total_values);
let mut new_offsets = Vec::with_capacity(row_count + 1);
Expand All @@ -420,6 +480,26 @@ fn array_sort_non_primitive<OffsetSize: OffsetSizeTrait>(
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;

Expand All @@ -428,7 +508,29 @@ fn array_sort_non_primitive<OffsetSize: OffsetSizeTrait>(
} 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)));
}

Expand Down
Loading