Skip to main content

datafusion_functions_nested/
sort.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! [`ScalarUDFImpl`] definitions for array_sort function.
19
20use 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/// Implementation of `array_sort` function
49///
50/// `array_sort` sorts the elements of an array
51///
52/// # Example
53///
54/// `array_sort([3, 1, 2])` returns `[1, 2, 3]`
55#[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        // Signature should prevent this arm ever occurring
200        _ => 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
218/// Sort each row of a primitive-typed ListArray using a custom in-place sort
219/// kernel.
220fn 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
248/// Fast path for primitive values with no element-level nulls. Copies all
249/// values into a single `Vec` and sorts each row's slice in-place.
250fn 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    // Copy all values into a mutable buffer
267    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
298/// Slow path for primitive values with element-level nulls.
299fn 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        // Compact valid values directly into the target region of the output
342        // buffer: after nulls (if nulls_first) or at the start (if nulls_last).
343        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        // Build validity bits
361        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
386/// Sort a non-pritive-typed ListArray by converting all rows at once using
387/// `RowConverter`, and then sort row indices by comparing encoded bytes (sort
388/// direction and null ordering are baked into the encoding), and materialize
389/// the result with a single `take()`.
390fn 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
452/// Select elements from `values` at the given `indices` using `compute::take`.
453/// We consume `indices` in order to avoid an intermediate copy.
454fn 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
474/// Rebase offsets so they start at 0. For non-sliced ListArrays (the common
475/// case) offsets already start at 0 and we can clone the Arc-backed buffer
476/// cheaply instead of allocating a new Vec.
477fn 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}