#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SortDirection {
Asc,
Desc,
}
impl SortDirection {
pub fn parse(value: &str) -> Self {
if value.trim().eq_ignore_ascii_case("DESC") {
SortDirection::Desc
} else {
SortDirection::Asc
}
}
pub fn to_catalog_string(self) -> &'static str {
match self {
SortDirection::Asc => "ASC",
SortDirection::Desc => "DESC",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum NullOrder {
NullsFirst,
NullsLast,
}
impl NullOrder {
pub fn parse(value: &str) -> Self {
if value.trim().eq_ignore_ascii_case("NULLS_FIRST") {
NullOrder::NullsFirst
} else {
NullOrder::NullsLast
}
}
pub fn to_catalog_string(self) -> &'static str {
match self {
NullOrder::NullsFirst => "NULLS_FIRST",
NullOrder::NullsLast => "NULLS_LAST",
}
}
pub fn nulls_first(self) -> bool {
matches!(self, NullOrder::NullsFirst)
}
}
pub const DUCKDB_DIALECT: &str = "duckdb";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SortField {
pub sort_key_index: i32,
pub expression: String,
pub dialect: String,
pub direction: SortDirection,
pub null_order: NullOrder,
}
impl SortField {
pub fn column(
sort_key_index: i32,
column: impl Into<String>,
direction: SortDirection,
null_order: NullOrder,
) -> Self {
SortField {
sort_key_index,
expression: column.into(),
dialect: DUCKDB_DIALECT.to_string(),
direction,
null_order,
}
}
pub fn column_candidate(&self) -> Option<String> {
parse_bare_column(&self.expression)
}
}
fn parse_bare_column(expr: &str) -> Option<String> {
let trimmed = expr.trim();
if trimmed.is_empty() {
return None;
}
if let Some(inner) = trimmed.strip_prefix('"').and_then(|s| s.strip_suffix('"'))
&& !inner.is_empty()
&& !inner.contains('"')
{
return Some(inner.to_string());
}
let mut chars = trimmed.chars();
let first = chars.next()?;
if !(first.is_ascii_alphabetic() || first == '_') {
return None;
}
if trimmed
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_')
{
Some(trimmed.to_string())
} else {
None
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SortSpec {
pub sort_id: i64,
pub fields: Vec<SortField>,
}
impl SortSpec {
pub fn is_producible(&self) -> bool {
!self.fields.is_empty() && self.fields.iter().all(|f| f.column_candidate().is_some())
}
pub fn producible_columns(&self) -> Option<Vec<(String, SortDirection, NullOrder)>> {
self.fields
.iter()
.filter(|f| f.dialect.eq_ignore_ascii_case("duckdb"))
.map(|f| f.column_candidate().map(|c| (c, f.direction, f.null_order)))
.collect()
}
pub fn from_rows(rows: Vec<(i64, i32, String, String, String, String)>) -> Option<SortSpec> {
let sort_id = rows.first()?.0;
let fields = rows
.into_iter()
.map(
|(_, sort_key_index, expression, dialect, sort_direction, null_order)| SortField {
sort_key_index,
expression,
dialect,
direction: SortDirection::parse(&sort_direction),
null_order: NullOrder::parse(&null_order),
},
)
.collect();
Some(SortSpec {
sort_id,
fields,
})
}
}
#[cfg(feature = "write")]
pub(crate) fn sort_batches_by_spec(
batches: Vec<arrow::record_batch::RecordBatch>,
data_schema: &arrow::datatypes::Schema,
sort_spec: Option<&SortSpec>,
) -> crate::Result<Vec<arrow::record_batch::RecordBatch>> {
use arrow::array::{ArrayRef, RecordBatch};
use std::sync::Arc;
let Some(keys) = sort_spec.and_then(|spec| spec.producible_columns()) else {
return Ok(batches);
};
if keys.is_empty() || batches.iter().all(|b| b.num_rows() == 0) {
return Ok(batches);
}
let mut resolved = Vec::with_capacity(keys.len());
for (name, direction, null_order) in &keys {
let Ok(index) = data_schema.index_of(name) else {
return Ok(batches);
};
resolved.push((
index,
arrow::compute::SortOptions {
descending: matches!(direction, SortDirection::Desc),
nulls_first: null_order.nulls_first(),
},
));
}
let full_schema = batches[0].schema();
let combined = arrow::compute::concat_batches(&full_schema, &batches)?;
let sort_columns: Vec<arrow::compute::SortColumn> = resolved
.iter()
.map(|(index, options)| arrow::compute::SortColumn {
values: Arc::clone(combined.column(*index)),
options: Some(*options),
})
.collect();
let indices = arrow::compute::lexsort_to_indices(&sort_columns, None)?;
let sorted_columns = combined
.columns()
.iter()
.map(|c| arrow::compute::take(c, &indices, None))
.collect::<std::result::Result<Vec<ArrayRef>, _>>()?;
let sorted_columns = sorted_columns
.iter()
.zip(full_schema.fields())
.map(|(column, field)| crate::column_rename::coerce_column(column, field.data_type()))
.collect::<datafusion::common::Result<Vec<_>>>()?;
Ok(vec![RecordBatch::try_new(full_schema, sorted_columns)?])
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn direction_roundtrip_and_case_insensitive() {
assert_eq!(SortDirection::parse("ASC"), SortDirection::Asc);
assert_eq!(SortDirection::parse("desc"), SortDirection::Desc);
assert_eq!(SortDirection::parse("DESC"), SortDirection::Desc);
assert_eq!(SortDirection::parse("whatever"), SortDirection::Asc);
assert_eq!(SortDirection::Asc.to_catalog_string(), "ASC");
assert_eq!(SortDirection::Desc.to_catalog_string(), "DESC");
}
#[test]
fn null_order_roundtrip_and_case_insensitive() {
assert_eq!(NullOrder::parse("NULLS_FIRST"), NullOrder::NullsFirst);
assert_eq!(NullOrder::parse("nulls_first"), NullOrder::NullsFirst);
assert_eq!(NullOrder::parse("NULLS_LAST"), NullOrder::NullsLast);
assert_eq!(NullOrder::parse("anything"), NullOrder::NullsLast);
assert_eq!(NullOrder::NullsFirst.to_catalog_string(), "NULLS_FIRST");
assert_eq!(NullOrder::NullsLast.to_catalog_string(), "NULLS_LAST");
assert!(NullOrder::NullsFirst.nulls_first());
assert!(!NullOrder::NullsLast.nulls_first());
}
#[test]
fn bare_column_expressions_are_producible() {
assert_eq!(parse_bare_column("ts"), Some("ts".to_string()));
assert_eq!(
parse_bare_column(" device_id "),
Some("device_id".to_string())
);
assert_eq!(parse_bare_column("_x1"), Some("_x1".to_string()));
assert_eq!(parse_bare_column("\"My Col\""), Some("My Col".to_string()));
}
#[test]
fn non_column_expressions_are_not_producible() {
for expr in ["", "date_trunc('day', ts)", "a + b", "t.ts", "1", "ts, device_id", "\"\""] {
assert_eq!(
parse_bare_column(expr),
None,
"expr {expr:?} should not be a bare column"
);
}
}
#[test]
fn producible_only_when_all_keys_are_columns() {
let ok = SortSpec {
sort_id: 1,
fields: vec![
SortField::column(0, "device_id", SortDirection::Asc, NullOrder::NullsLast),
SortField::column(1, "ts", SortDirection::Desc, NullOrder::NullsFirst),
],
};
assert!(ok.is_producible());
assert_eq!(
ok.producible_columns().unwrap(),
vec![
(
"device_id".to_string(),
SortDirection::Asc,
NullOrder::NullsLast
),
("ts".to_string(), SortDirection::Desc, NullOrder::NullsFirst),
]
);
let mixed = SortSpec {
sort_id: 2,
fields: vec![
SortField::column(0, "device_id", SortDirection::Asc, NullOrder::NullsLast),
SortField {
sort_key_index: 1,
expression: "date_trunc('day', ts)".to_string(),
dialect: DUCKDB_DIALECT.to_string(),
direction: SortDirection::Asc,
null_order: NullOrder::NullsLast,
},
],
};
assert!(!mixed.is_producible());
assert_eq!(mixed.producible_columns(), None);
}
#[test]
fn from_rows_orders_and_parses() {
let rows = vec![
(
7,
0,
"device_id".to_string(),
"duckdb".to_string(),
"ASC".to_string(),
"NULLS_LAST".to_string(),
),
(
7,
1,
"ts".to_string(),
"duckdb".to_string(),
"DESC".to_string(),
"NULLS_FIRST".to_string(),
),
];
let spec = SortSpec::from_rows(rows).unwrap();
assert_eq!(spec.sort_id, 7);
assert_eq!(spec.fields.len(), 2);
assert_eq!(spec.fields[0].direction, SortDirection::Asc);
assert_eq!(spec.fields[0].null_order, NullOrder::NullsLast);
assert_eq!(spec.fields[1].direction, SortDirection::Desc);
assert_eq!(spec.fields[1].null_order, NullOrder::NullsFirst);
assert!(spec.is_producible());
}
#[test]
fn from_rows_empty_is_none() {
assert_eq!(SortSpec::from_rows(vec![]), None);
}
#[cfg(feature = "write")]
#[test]
fn sort_batches_preserves_nested_field_metadata() {
use std::{collections::HashMap, sync::Arc};
use arrow::{
array::{ArrayRef, Int32Array, Int64Array, StringArray, StructArray},
datatypes::{DataType, Field, Schema},
record_batch::RecordBatch,
};
let field_id =
|value: &str| HashMap::from([("PARQUET:field_id".to_string(), value.to_string())]);
let money_fields = vec![
Arc::new(Field::new("amount", DataType::Int32, false).with_metadata(field_id("2"))),
Arc::new(Field::new("currency", DataType::Utf8, false).with_metadata(field_id("3"))),
];
let money: ArrayRef = Arc::new(StructArray::new(
money_fields.clone().into(),
vec![
Arc::new(Int32Array::from(vec![20, 10])),
Arc::new(StringArray::from(vec!["EUR", "USD"])),
],
None,
));
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int64, false),
Field::new("money", DataType::Struct(money_fields.into()), false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int64Array::from(vec![2, 1])), money],
)
.unwrap();
let sort_spec = SortSpec {
sort_id: 1,
fields: vec![SortField::column(0, "id", SortDirection::Asc, NullOrder::NullsLast)],
};
let sorted = sort_batches_by_spec(vec![batch], &schema, Some(&sort_spec)).unwrap();
assert_eq!(
sorted[0]
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap()
.values(),
&[1, 2]
);
assert_eq!(
sorted[0].schema().field(1).data_type(),
sorted[0].column(1).data_type()
);
let DataType::Struct(fields) = sorted[0].column(1).data_type() else {
panic!("money must remain a struct");
};
assert_eq!(
fields[0].metadata().get("PARQUET:field_id"),
Some(&"2".into())
);
assert_eq!(
fields[1].metadata().get("PARQUET:field_id"),
Some(&"3".into())
);
}
}