Skip to content
Open
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
32 changes: 26 additions & 6 deletions datafusion/common/src/utils/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ use arrow::compute::kernels::cmp::eq;
use arrow::compute::kernels::length::length;
use arrow::compute::{SortColumn, SortOptions, nullif, partition};
use arrow::datatypes::{
ArrowNativeType, DataType, Field, Int32Type, Int64Type, SchemaRef,
ArrowNativeType, DataType, Field, Int32Type, Int64Type, Metadata, SchemaRef,
};
#[cfg(feature = "sql")]
use sqlparser::{ast::Ident, dialect::GenericDialect, parser::Parser};
Expand Down Expand Up @@ -559,6 +559,7 @@ pub struct SingleRowListArrayBuilder {
/// Specify the field name for the resulting array. Defaults to value used in
/// [`Field::new_list_field`]
field_name: Option<String>,
field_metadata: Metadata,
}

impl SingleRowListArrayBuilder {
Expand All @@ -568,6 +569,7 @@ impl SingleRowListArrayBuilder {
arr,
nullable: true,
field_name: None,
field_metadata: Metadata::default(),
}
}

Expand All @@ -583,10 +585,12 @@ impl SingleRowListArrayBuilder {
self
}

/// Copies field name and nullable from the specified field
pub fn with_field(self, field: &Field) -> Self {
self.with_field_name(Some(field.name().to_owned()))
.with_nullable(field.is_nullable())
/// Copies field name, nullability, and metadata from the specified field
pub fn with_field(mut self, field: &Field) -> Self {
self.field_name = Some(field.name().to_owned());
self.nullable = field.is_nullable();
self.field_metadata = field.metadata().clone();
self
}

/// Build a single element [`ListArray`]
Expand Down Expand Up @@ -659,13 +663,14 @@ impl SingleRowListArrayBuilder {
arr,
nullable,
field_name,
field_metadata,
} = self;
let data_type = arr.data_type().to_owned();
let field = match field_name {
Some(name) => Field::new(name, data_type, nullable),
None => Field::new_list_field(data_type, nullable),
};
(Arc::new(field), arr)
(Arc::new(field.with_metadata(field_metadata)), arr)
}
}

Expand Down Expand Up @@ -1606,6 +1611,7 @@ pub fn apply_parent_nulls(

#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::sync::Arc;

use super::*;
Expand All @@ -1617,6 +1623,20 @@ mod tests {
#[cfg(feature = "sql")]
use sqlparser::ast::Ident;

#[test]
fn test_list_scalar_cast_preserves_target_element_metadata() -> Result<()> {
let input =
ListArray::from_iter_primitive::<Int32Type, _, _>(vec![Some(vec![Some(1)])]);
let item = Field::new("item", DataType::Int32, true)
.with_metadata(HashMap::from([("PARQUET:field_id".into(), "2".into())]));
let target = DataType::List(Arc::new(item));

let casted = ScalarValue::List(Arc::new(input)).cast_to(&target)?;
assert_eq!(casted.data_type(), target);
assert_eq!(casted.to_array_of_size(2)?.data_type(), &target);
Ok(())
}

#[test]
fn test_apply_parent_nulls_sliced() -> Result<()> {
let child = Arc::new(
Expand Down
Loading