1use crate::utils::make_scalar_function;
21use arrow::array::BooleanBufferBuilder;
22use arrow::array::{
23 Array, ArrayRef, ArrowPrimitiveType, GenericListArray, OffsetSizeTrait,
24 PrimitiveArray, UInt32Array, UInt64Array, new_empty_array, new_null_array,
25};
26use arrow::buffer::{NullBuffer, OffsetBuffer};
27use arrow::datatypes::{ArrowNativeTypeOp, DataType, FieldRef};
28use arrow::row::{RowConverter, SortField};
29use arrow::{compute, compute::SortOptions, downcast_primitive_array};
30use datafusion_common::cast::{as_large_list_array, as_list_array, as_string_array};
31use datafusion_common::utils::ListCoercion;
32use datafusion_common::{Result, exec_err, internal_datafusion_err};
33use datafusion_expr::{
34 ArrayFunctionArgument, ArrayFunctionSignature, ColumnarValue, Documentation,
35 ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignature, Volatility,
36};
37use datafusion_macros::user_doc;
38use std::sync::Arc;
39
40make_udf_expr_and_func!(
41 ArraySort,
42 array_sort,
43 array desc null_first,
44 "returns sorted array.",
45 array_sort_udf
46);
47
48#[user_doc(
56 doc_section(label = "Array Functions"),
57 description = "Sort array.",
58 syntax_example = "array_sort(array[, order[, nulls_order]])",
59 sql_example = r#"```sql
60> select array_sort([3, 1, 2]);
61+-----------------------------+
62| array_sort(List([3,1,2])) |
63+-----------------------------+
64| [1, 2, 3] |
65+-----------------------------+
66> select array_sort([3, 1, NULL, 2], 'desc', 'nulls last');
67+--------------------------------------------------+
68| array_sort(List(3,1,NULL,2),'desc','nulls last') |
69+--------------------------------------------------+
70| [3, 2, 1, NULL] |
71+--------------------------------------------------+
72```"#,
73 argument(
74 name = "array",
75 description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
76 ),
77 argument(
78 name = "order",
79 description = "Whether to sort in ascending (`ASC`) or descending (`DESC`) order. The default is `ASC`."
80 ),
81 argument(
82 name = "nulls_order",
83 description = "Whether to sort nulls first (`NULLS FIRST`) or last (`NULLS LAST`). The default is `NULLS FIRST`."
84 )
85)]
86#[derive(Debug, PartialEq, Eq, Hash)]
87pub struct ArraySort {
88 signature: Signature,
89 aliases: Vec<String>,
90}
91
92impl Default for ArraySort {
93 fn default() -> Self {
94 Self::new()
95 }
96}
97
98impl ArraySort {
99 pub fn new() -> Self {
100 Self {
101 signature: Signature::one_of(
102 vec![
103 TypeSignature::ArraySignature(ArrayFunctionSignature::Array {
104 arguments: vec![ArrayFunctionArgument::Array],
105 array_coercion: Some(ListCoercion::FixedSizedListToList),
106 }),
107 TypeSignature::ArraySignature(ArrayFunctionSignature::Array {
108 arguments: vec![
109 ArrayFunctionArgument::Array,
110 ArrayFunctionArgument::String,
111 ],
112 array_coercion: Some(ListCoercion::FixedSizedListToList),
113 }),
114 TypeSignature::ArraySignature(ArrayFunctionSignature::Array {
115 arguments: vec![
116 ArrayFunctionArgument::Array,
117 ArrayFunctionArgument::String,
118 ArrayFunctionArgument::String,
119 ],
120 array_coercion: Some(ListCoercion::FixedSizedListToList),
121 }),
122 ],
123 Volatility::Immutable,
124 ),
125 aliases: vec!["list_sort".to_string()],
126 }
127 }
128}
129
130impl ScalarUDFImpl for ArraySort {
131 fn name(&self) -> &str {
132 "array_sort"
133 }
134
135 fn signature(&self) -> &Signature {
136 &self.signature
137 }
138
139 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
140 Ok(arg_types[0].clone())
141 }
142
143 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
144 make_scalar_function(array_sort_inner)(&args.args)
145 }
146
147 fn aliases(&self) -> &[String] {
148 &self.aliases
149 }
150
151 fn documentation(&self) -> Option<&Documentation> {
152 self.doc()
153 }
154}
155
156fn array_sort_inner(args: &[ArrayRef]) -> Result<ArrayRef> {
157 if args.is_empty() || args.len() > 3 {
158 return exec_err!("array_sort expects one to three arguments");
159 }
160
161 if args[0].is_empty() || args[0].data_type().is_null() {
162 return Ok(Arc::clone(&args[0]));
163 }
164
165 if args[1..].iter().any(|array| array.is_null(0)) {
166 return Ok(new_null_array(args[0].data_type(), args[0].len()));
167 }
168
169 let sort_options = if args.len() >= 2 {
170 let order = as_string_array(&args[1])?.value(0);
171 let descending = order_desc(order)?;
172 let nulls_first = if args.len() >= 3 {
173 order_nulls_first(as_string_array(&args[2])?.value(0))?
174 } else {
175 true
176 };
177 Some(SortOptions {
178 descending,
179 nulls_first,
180 })
181 } else {
182 None
183 };
184
185 match args[0].data_type() {
186 DataType::List(field) | DataType::LargeList(field)
187 if field.data_type().is_null() =>
188 {
189 Ok(Arc::clone(&args[0]))
190 }
191 DataType::List(field) => {
192 let array = as_list_array(&args[0])?;
193 array_sort_generic(array, Arc::clone(field), sort_options)
194 }
195 DataType::LargeList(field) => {
196 let array = as_large_list_array(&args[0])?;
197 array_sort_generic(array, Arc::clone(field), sort_options)
198 }
199 _ => exec_err!("array_sort expects list for first argument"),
201 }
202}
203
204fn array_sort_generic<OffsetSize: OffsetSizeTrait>(
205 list_array: &GenericListArray<OffsetSize>,
206 field: FieldRef,
207 sort_options: Option<SortOptions>,
208) -> Result<ArrayRef> {
209 let values = list_array.values();
210
211 if values.data_type().is_primitive() {
212 array_sort_primitive(list_array, field, sort_options)
213 } else {
214 array_sort_non_primitive(list_array, field, sort_options)
215 }
216}
217
218fn array_sort_primitive<OffsetSize: OffsetSizeTrait>(
221 list_array: &GenericListArray<OffsetSize>,
222 field: FieldRef,
223 sort_options: Option<SortOptions>,
224) -> Result<ArrayRef> {
225 let values = list_array.values().as_ref();
226 downcast_primitive_array! {
227 values => sort_primitive_list(values, list_array, field, sort_options),
228 _ => exec_err!("array_sort: unsupported primitive type")
229 }
230}
231
232fn sort_primitive_list<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
233 prim_values: &PrimitiveArray<T>,
234 list_array: &GenericListArray<OffsetSize>,
235 field: FieldRef,
236 sort_options: Option<SortOptions>,
237) -> Result<ArrayRef>
238where
239 T::Native: ArrowNativeTypeOp,
240{
241 if prim_values.null_count() > 0 {
242 sort_list_with_nulls(prim_values, list_array, field, sort_options)
243 } else {
244 sort_list_no_nulls(prim_values, list_array, field, sort_options)
245 }
246}
247
248fn sort_list_no_nulls<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
251 prim_values: &PrimitiveArray<T>,
252 list_array: &GenericListArray<OffsetSize>,
253 field: FieldRef,
254 sort_options: Option<SortOptions>,
255) -> Result<ArrayRef>
256where
257 T::Native: ArrowNativeTypeOp,
258{
259 let row_count = list_array.len();
260 let offsets = list_array.offsets();
261 let values_start = offsets[0].as_usize();
262 let values_end = offsets[row_count].as_usize();
263
264 let descending = sort_options.is_some_and(|o| o.descending);
265
266 let mut values: Vec<T::Native> =
268 prim_values.values()[values_start..values_end].to_vec();
269
270 for (row_index, window) in offsets.windows(2).enumerate() {
271 if list_array.is_null(row_index) {
272 continue;
273 }
274 let start = window[0].as_usize() - values_start;
275 let end = window[1].as_usize() - values_start;
276 let slice = &mut values[start..end];
277 if descending {
278 slice.sort_unstable_by(|a, b| b.compare(*a));
279 } else {
280 slice.sort_unstable_by(|a, b| a.compare(*b));
281 }
282 }
283
284 let new_offsets = rebase_offsets(offsets);
285 let sorted_values = Arc::new(
286 PrimitiveArray::<T>::new(values.into(), None)
287 .with_data_type(prim_values.data_type().clone()),
288 );
289
290 Ok(Arc::new(GenericListArray::<OffsetSize>::try_new(
291 field,
292 new_offsets,
293 sorted_values,
294 list_array.nulls().cloned(),
295 )?))
296}
297
298fn sort_list_with_nulls<T: ArrowPrimitiveType, OffsetSize: OffsetSizeTrait>(
300 prim_values: &PrimitiveArray<T>,
301 list_array: &GenericListArray<OffsetSize>,
302 field: FieldRef,
303 sort_options: Option<SortOptions>,
304) -> Result<ArrayRef>
305where
306 T::Native: ArrowNativeTypeOp,
307{
308 let row_count = list_array.len();
309 let offsets = list_array.offsets();
310 let values_start = offsets[0].as_usize();
311 let values_end = offsets[row_count].as_usize();
312 let total_values = values_end - values_start;
313
314 let descending = sort_options.is_some_and(|o| o.descending);
315 let nulls_first = sort_options.is_none_or(|o| o.nulls_first);
316
317 let mut out_values: Vec<T::Native> = vec![T::Native::default(); total_values];
318 let mut validity = BooleanBufferBuilder::new(total_values);
319
320 let src_nulls = prim_values.nulls().ok_or_else(|| {
321 internal_datafusion_err!(
322 "sort_list_with_nulls called but values have no null buffer"
323 )
324 })?;
325 let src_values = prim_values.values();
326
327 for (row_index, window) in offsets.windows(2).enumerate() {
328 let start = window[0].as_usize();
329 let end = window[1].as_usize();
330 let row_len = end - start;
331 let out_start = start - values_start;
332
333 if list_array.is_null(row_index) || row_len == 0 {
334 validity.append_n(row_len, false);
335 continue;
336 }
337
338 let null_count = src_nulls.slice(start, row_len).null_count();
339 let valid_count = row_len - null_count;
340
341 let valid_offset = if nulls_first { null_count } else { 0 };
344 let mut write_pos = out_start + valid_offset;
345 for i in start..end {
346 if src_nulls.is_valid(i) {
347 out_values[write_pos] = src_values[i];
348 write_pos += 1;
349 }
350 }
351
352 let valid_slice = &mut out_values
353 [out_start + valid_offset..out_start + valid_offset + valid_count];
354 if descending {
355 valid_slice.sort_unstable_by(|a, b| b.compare(*a));
356 } else {
357 valid_slice.sort_unstable_by(|a, b| a.compare(*b));
358 }
359
360 if nulls_first {
362 validity.append_n(null_count, false);
363 validity.append_n(valid_count, true);
364 } else {
365 validity.append_n(valid_count, true);
366 validity.append_n(null_count, false);
367 }
368 }
369
370 let new_offsets = rebase_offsets(offsets);
371
372 let null_buffer = NullBuffer::from(validity.finish());
373 let sorted_values = Arc::new(
374 PrimitiveArray::<T>::new(out_values.into(), Some(null_buffer))
375 .with_data_type(prim_values.data_type().clone()),
376 );
377
378 Ok(Arc::new(GenericListArray::<OffsetSize>::try_new(
379 field,
380 new_offsets,
381 sorted_values,
382 list_array.nulls().cloned(),
383 )?))
384}
385
386fn array_sort_non_primitive<OffsetSize: OffsetSizeTrait>(
391 list_array: &GenericListArray<OffsetSize>,
392 field: FieldRef,
393 sort_options: Option<SortOptions>,
394) -> Result<ArrayRef> {
395 let row_count = list_array.len();
396 let values = list_array.values();
397 let offsets = list_array.offsets();
398 let values_start = offsets[0].as_usize();
399 let total_values = offsets[row_count].as_usize() - values_start;
400
401 let converter = RowConverter::new(vec![SortField::new_with_options(
402 values.data_type().clone(),
403 sort_options.unwrap_or_default(),
404 )])?;
405 let values_sliced = values.slice(values_start, total_values);
406 let rows = converter.convert_columns(&[Arc::clone(&values_sliced)])?;
407
408 let mut indices: Vec<OffsetSize> = Vec::with_capacity(total_values);
409 let mut new_offsets = Vec::with_capacity(row_count + 1);
410 new_offsets.push(OffsetSize::usize_as(0));
411
412 let mut sort_scratch: Vec<usize> = Vec::new();
413
414 for (row_index, window) in offsets.windows(2).enumerate() {
415 let start = window[0];
416 let end = window[1];
417
418 if list_array.is_null(row_index) {
419 new_offsets.push(new_offsets[row_index]);
420 continue;
421 }
422
423 let len = (end - start).as_usize();
424 let local_start = start.as_usize() - values_start;
425
426 if len <= 1 {
427 indices.extend((local_start..local_start + len).map(OffsetSize::usize_as));
428 } else {
429 sort_scratch.clear();
430 sort_scratch.extend(local_start..local_start + len);
431 sort_scratch.sort_unstable_by(|&a, &b| rows.row(a).cmp(&rows.row(b)));
432 indices.extend(sort_scratch.iter().map(|&i| OffsetSize::usize_as(i)));
433 }
434
435 new_offsets.push(new_offsets[row_index] + (end - start));
436 }
437
438 let sorted_values = if indices.is_empty() {
439 new_empty_array(values.data_type())
440 } else {
441 take_by_indices(&values_sliced, indices)?
442 };
443
444 Ok(Arc::new(GenericListArray::<OffsetSize>::try_new(
445 field,
446 OffsetBuffer::<OffsetSize>::new(new_offsets.into()),
447 sorted_values,
448 list_array.nulls().cloned(),
449 )?))
450}
451
452fn take_by_indices<OffsetSize: OffsetSizeTrait>(
455 values: &ArrayRef,
456 indices: Vec<OffsetSize>,
457) -> Result<ArrayRef> {
458 let len = indices.len();
459 let buffer = arrow::buffer::Buffer::from_vec(indices);
460 let indices_array: ArrayRef = if OffsetSize::IS_LARGE {
461 Arc::new(UInt64Array::new(
462 arrow::buffer::ScalarBuffer::new(buffer, 0, len),
463 None,
464 ))
465 } else {
466 Arc::new(UInt32Array::new(
467 arrow::buffer::ScalarBuffer::new(buffer, 0, len),
468 None,
469 ))
470 };
471 Ok(compute::take(values.as_ref(), &indices_array, None)?)
472}
473
474fn rebase_offsets<OffsetSize: OffsetSizeTrait>(
478 offsets: &OffsetBuffer<OffsetSize>,
479) -> OffsetBuffer<OffsetSize> {
480 offsets.clone().subtract(offsets[0])
481}
482
483fn order_desc(modifier: &str) -> Result<bool> {
484 match modifier.to_uppercase().as_str() {
485 "DESC" => Ok(true),
486 "ASC" => Ok(false),
487 _ => exec_err!("the second parameter of array_sort expects DESC or ASC"),
488 }
489}
490
491fn order_nulls_first(modifier: &str) -> Result<bool> {
492 match modifier.to_uppercase().as_str() {
493 "NULLS FIRST" => Ok(true),
494 "NULLS LAST" => Ok(false),
495 _ => exec_err!(
496 "the third parameter of array_sort expects NULLS FIRST or NULLS LAST"
497 ),
498 }
499}