use crate::Result;
use crate::error::CoreError;
use arrow::array::ArrayRef;
use arrow::array::RecordBatch;
use arrow::array::StringArray;
use arrow_array::{Array, UInt32Array};
use arrow_row::{RowConverter, SortField};
use arrow_schema::{ArrowError, SchemaRef};
pub trait ColumnAsArray {
fn get_array(&self, column_name: &str) -> Result<ArrayRef>;
fn get_string_array(&self, column_name: &str) -> Result<StringArray>;
}
impl ColumnAsArray for RecordBatch {
fn get_array(&self, column_name: &str) -> Result<ArrayRef> {
let index = self.schema().index_of(column_name)?;
let array = self.column(index);
Ok(array.clone())
}
fn get_string_array(&self, column_name: &str) -> Result<StringArray> {
let array = self.get_array(column_name)?;
let array = array
.as_any()
.downcast_ref::<StringArray>()
.ok_or_else(|| {
ArrowError::CastError(format!(
"Column {column_name} cannot be cast to StringArray."
))
})?;
Ok(array.clone())
}
}
pub fn lexsort_to_indices(arrays: &[ArrayRef], desc: bool) -> UInt32Array {
let fields = arrays
.iter()
.map(|a| SortField::new(a.data_type().clone()))
.collect();
let converter = RowConverter::new(fields).unwrap();
let rows = converter.convert_columns(arrays).unwrap();
let mut sort: Vec<_> = rows.iter().enumerate().collect();
if desc {
sort.sort_unstable_by(|(ia, a), (ib, b)| b.cmp(a).then(ia.cmp(ib)));
} else {
sort.sort_unstable_by(|(ia, a), (ib, b)| a.cmp(b).then(ia.cmp(ib)));
}
UInt32Array::from_iter_values(sort.iter().map(|(i, _)| *i as u32))
}
pub fn create_row_converter<I, S>(schema: SchemaRef, column_names: I) -> Result<RowConverter>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let sort_fields: Result<Vec<_>> = column_names
.into_iter()
.map(|col| {
let (_, field) = schema
.column_with_name(col.as_ref())
.ok_or_else(|| CoreError::Schema(format!("Column {} not found", col.as_ref())))?;
Ok(SortField::new(field.data_type().clone()))
})
.collect();
RowConverter::new(sort_fields?).map_err(CoreError::ArrowError)
}
pub fn get_column_arrays<I, S>(batch: &RecordBatch, column_names: I) -> Result<Vec<ArrayRef>>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
column_names
.into_iter()
.map(|col| {
batch
.column_by_name(col.as_ref())
.cloned()
.ok_or_else(|| CoreError::Schema(format!("Column {} not found", col.as_ref())))
})
.collect()
}
pub fn project_batch_by_names(
batch: RecordBatch,
projection: Option<&[String]>,
) -> Result<RecordBatch> {
let Some(cols) = projection else {
return Ok(batch);
};
let indices: Vec<usize> = cols
.iter()
.map(|name| {
batch
.schema()
.index_of(name)
.map_err(|e| CoreError::Schema(format!("Projection column not found: {e:?}")))
})
.collect::<Result<_>>()?;
batch.project(&indices).map_err(CoreError::ArrowError)
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{Int32Array, StringArray};
use arrow_array::Float64Array;
use std::sync::Arc;
#[test]
fn test_basic_int_sort() {
let arr = Int32Array::from(vec![3, 1, 4, 1, 5]);
let arrays = vec![Arc::new(arr) as ArrayRef];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(
result.values(),
&[1, 3, 0, 2, 4] );
let result = lexsort_to_indices(&arrays, true);
assert_eq!(
result.values(),
&[4, 2, 0, 1, 3] );
}
#[test]
fn test_multiple_columns() {
let arr1 = Int32Array::from(vec![1, 1, 2, 2]);
let arr2 = StringArray::from(vec!["b", "a", "b", "a"]);
let arrays = vec![Arc::new(arr1) as ArrayRef, Arc::new(arr2) as ArrayRef];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(
result.values(),
&[1, 0, 3, 2] );
}
#[test]
fn test_edge_cases() {
assert_eq!(lexsort_to_indices(&[], false).len(), 0);
let arr = Int32Array::from(vec![] as Vec<i32>);
let arrays = vec![Arc::new(arr) as ArrayRef];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(result.len(), 0);
let arr = Int32Array::from(vec![1]);
let arrays = vec![Arc::new(arr) as ArrayRef];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(result.values(), &[0]);
let arr = Int32Array::from(vec![5, 5, 5, 5]);
let arrays = vec![Arc::new(arr) as ArrayRef];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(result.values(), &[0, 1, 2, 3]);
}
#[test]
fn test_different_types() {
let int_arr = Int32Array::from(vec![1, 2, 1]);
let str_arr = StringArray::from(vec!["a", "b", "c"]);
let float_arr = Float64Array::from(vec![1.0, 2.0, 3.0]);
let arrays = vec![
Arc::new(int_arr) as ArrayRef,
Arc::new(str_arr) as ArrayRef,
Arc::new(float_arr) as ArrayRef,
];
let result = lexsort_to_indices(&arrays, false);
assert_eq!(result.values(), &[0, 2, 1]);
}
#[test]
fn test_project_batch_by_names() {
use arrow_schema::{DataType, Field, Schema};
let schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Int32, false),
Field::new("b", DataType::Utf8, false),
Field::new("c", DataType::Float64, false),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1, 2])) as ArrayRef,
Arc::new(StringArray::from(vec!["x", "y"])) as ArrayRef,
Arc::new(Float64Array::from(vec![1.5, 2.5])) as ArrayRef,
],
)
.unwrap();
let same = project_batch_by_names(batch.clone(), None).unwrap();
assert_eq!(same.num_columns(), 3);
assert_eq!(same.schema().field(0).name(), "a");
let cols = vec!["c".to_string(), "a".to_string()];
let projected = project_batch_by_names(batch.clone(), Some(&cols)).unwrap();
assert_eq!(projected.num_columns(), 2);
assert_eq!(projected.schema().field(0).name(), "c");
assert_eq!(projected.schema().field(1).name(), "a");
let bad = vec!["a".to_string(), "missing".to_string()];
let err = project_batch_by_names(batch, Some(&bad)).unwrap_err();
assert!(matches!(err, CoreError::Schema(_)));
assert!(err.to_string().contains("missing"));
}
}