use std::collections::{BTreeMap, BTreeSet};
use std::sync::Arc;
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
use datafusion_functions::core::input_file_name::InputFileNameFunc;
use parquet::arrow::ProjectionMask;
use parquet::schema::types::SchemaDescriptor;
use datafusion_common::Result;
use datafusion_common::nested_struct::requires_nested_struct_cast;
use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion, TreeNodeVisitor};
use datafusion_functions::core::file_row_index::FileRowIndexFunc;
use datafusion_functions::core::getfield::GetFieldFunc;
use datafusion_physical_expr::expressions::{CastExpr, Column, Literal};
use datafusion_physical_expr::utils::collect_columns;
use datafusion_physical_expr::{PhysicalExpr, ScalarFunctionExpr};
use crate::nested_schema_pruning::{
CastColumnAccess, clip_for_cast, contains_struct, count_leaves, field_with_type,
};
#[derive(Debug, Clone)]
pub(crate) struct ParquetReadPlan {
pub projection_mask: ProjectionMask,
pub projected_schema: SchemaRef,
}
#[derive(Debug, Clone)]
pub(crate) struct StructFieldAccess {
pub(crate) root_index: usize,
pub(crate) field_path: Vec<String>,
}
#[derive(Debug, Default)]
struct StructAccessTree<'a> {
roots: BTreeMap<usize, StructAccessNode<'a>>,
}
#[derive(Debug, Default)]
struct StructAccessNode<'a> {
children: BTreeMap<&'a str, StructAccessNode<'a>>,
selected_here: bool,
}
impl<'a> StructAccessTree<'a> {
fn from_accesses(accesses: &'a [StructFieldAccess]) -> Self {
let mut tree = Self::default();
for StructFieldAccess {
root_index,
field_path,
} in accesses
{
let mut node = tree.roots.entry(*root_index).or_default();
for component in field_path {
node = node.children.entry(component.as_str()).or_default();
}
node.selected_here = true;
}
tree
}
fn root(&self, idx: usize) -> Option<&StructAccessNode<'a>> {
self.roots.get(&idx)
}
}
pub(crate) struct PushdownChecker<'schema> {
non_primitive_columns: bool,
projected_columns: bool,
has_unpushable_udfs: bool,
required_columns: Vec<usize>,
struct_field_accesses: Vec<StructFieldAccess>,
cast_accesses: Vec<CastColumnAccess>,
collect_cast_accesses: bool,
allow_list_columns: bool,
file_schema: &'schema Schema,
}
impl<'schema> PushdownChecker<'schema> {
pub(crate) fn new(file_schema: &'schema Schema, allow_list_columns: bool) -> Self {
Self {
non_primitive_columns: false,
projected_columns: false,
has_unpushable_udfs: false,
required_columns: Vec::new(),
struct_field_accesses: Vec::new(),
cast_accesses: Vec::new(),
collect_cast_accesses: false,
allow_list_columns,
file_schema,
}
}
pub(crate) fn with_cast_collection(mut self) -> Self {
self.collect_cast_accesses = true;
self
}
fn check_struct_field_column(
&mut self,
column_name: &str,
field_path: Vec<String>,
) -> Option<TreeNodeRecursion> {
let Ok(idx) = self.file_schema.index_of(column_name) else {
self.projected_columns = true;
return Some(TreeNodeRecursion::Jump);
};
self.struct_field_accesses.push(StructFieldAccess {
root_index: idx,
field_path,
});
None
}
fn check_single_column(&mut self, column_name: &str) -> Option<TreeNodeRecursion> {
let idx = match self.file_schema.index_of(column_name) {
Ok(idx) => idx,
Err(_) => {
self.projected_columns = true;
return Some(TreeNodeRecursion::Jump);
}
};
self.required_columns.push(idx);
let data_type = self.file_schema.field(idx).data_type();
if DataType::is_nested(data_type) {
self.handle_nested_type(data_type)
} else {
None
}
}
fn handle_nested_type(&mut self, data_type: &DataType) -> Option<TreeNodeRecursion> {
if self.is_nested_type_supported(data_type) {
None
} else {
self.non_primitive_columns = true;
Some(TreeNodeRecursion::Jump)
}
}
fn is_nested_type_supported(&self, data_type: &DataType) -> bool {
let is_list = matches!(
data_type,
DataType::List(_) | DataType::LargeList(_) | DataType::FixedSizeList(_, _)
);
self.allow_list_columns && is_list
}
#[inline]
pub(crate) fn prevents_pushdown(&self) -> bool {
self.non_primitive_columns || self.projected_columns || self.has_unpushable_udfs
}
pub(crate) fn into_sorted_columns(mut self) -> PushdownColumns {
self.required_columns.sort_unstable();
self.required_columns.dedup();
PushdownColumns {
required_columns: self.required_columns,
struct_field_accesses: self.struct_field_accesses,
cast_accesses: self.cast_accesses,
}
}
}
impl TreeNodeVisitor<'_> for PushdownChecker<'_> {
type Node = Arc<dyn PhysicalExpr>;
fn f_down(&mut self, node: &Self::Node) -> Result<TreeNodeRecursion> {
if let Some(func) =
ScalarFunctionExpr::try_downcast_func::<GetFieldFunc>(node.as_ref())
{
let args = func.args();
if let Some(column) = args.first().and_then(|a| a.downcast_ref::<Column>()) {
let is_map_column = self
.file_schema
.index_of(column.name())
.ok()
.map(|idx| {
matches!(
self.file_schema.field(idx).data_type(),
DataType::Map(_, _)
)
})
.unwrap_or(false);
let return_type = func.return_type();
if !is_map_column
&& (!DataType::is_nested(return_type)
|| self.is_nested_type_supported(return_type))
{
let field_path = args[1..]
.iter()
.map(|arg| {
arg.downcast_ref::<Literal>().and_then(|lit| {
lit.value().try_as_str().flatten().map(|s| s.to_string())
})
})
.collect();
match field_path {
Some(path) => {
if let Some(recursion) =
self.check_struct_field_column(column.name(), path)
{
return Ok(recursion);
}
}
None => {
if let Some(recursion) =
self.check_single_column(column.name())
{
return Ok(recursion);
}
}
}
return Ok(TreeNodeRecursion::Jump);
}
}
}
if self.collect_cast_accesses
&& let Some(cast) = node.downcast_ref::<CastExpr>()
&& let Some(column) = cast.expr().downcast_ref::<Column>()
&& let Ok(idx) = self.file_schema.index_of(column.name())
&& requires_nested_struct_cast(
self.file_schema.field(idx).data_type(),
cast.cast_type(),
)
{
self.cast_accesses.push(CastColumnAccess {
root_index: idx,
target_type: cast.cast_type().clone(),
});
return Ok(TreeNodeRecursion::Jump);
}
if let Some(column) = node.downcast_ref::<Column>()
&& let Some(recursion) = self.check_single_column(column.name())
{
return Ok(recursion);
}
if ScalarFunctionExpr::try_downcast_func::<InputFileNameFunc>(node.as_ref())
.is_some()
|| ScalarFunctionExpr::try_downcast_func::<FileRowIndexFunc>(node.as_ref())
.is_some()
{
self.has_unpushable_udfs = true;
return Ok(TreeNodeRecursion::Jump);
}
Ok(TreeNodeRecursion::Continue)
}
}
#[derive(Debug)]
pub(crate) struct PushdownColumns {
pub(crate) required_columns: Vec<usize>,
pub(crate) struct_field_accesses: Vec<StructFieldAccess>,
pub(crate) cast_accesses: Vec<CastColumnAccess>,
}
pub(crate) fn build_projection_read_plan(
exprs: impl IntoIterator<Item = Arc<dyn PhysicalExpr>>,
file_schema: &Schema,
schema_descr: &SchemaDescriptor,
) -> ParquetReadPlan {
let exprs = exprs.into_iter().collect::<Vec<_>>();
let all_plain_columns = exprs.iter().all(|e| e.downcast_ref::<Column>().is_some());
if all_plain_columns {
let mut root_indices: Vec<usize> = exprs
.iter()
.map(|e| e.downcast_ref::<Column>().unwrap().index())
.collect();
root_indices.sort_unstable();
root_indices.dedup();
return root_level_plan(&root_indices, file_schema, schema_descr);
}
let projected_columns = exprs.iter().flat_map(collect_columns).collect::<Vec<_>>();
let all_resolvable_and_struct_free = projected_columns.iter().all(|col| {
file_schema
.fields()
.get(col.index())
.is_some_and(|f| f.name() == col.name() && !contains_struct(f.data_type()))
});
if all_resolvable_and_struct_free {
let mut root_indices = projected_columns
.iter()
.map(|c| c.index())
.collect::<Vec<_>>();
root_indices.sort_unstable();
root_indices.dedup();
return root_level_plan(&root_indices, file_schema, schema_descr);
}
let mut all_root_indices = Vec::new();
let mut all_struct_accesses = Vec::new();
let mut all_cast_accesses = Vec::new();
for expr in exprs {
let mut checker = PushdownChecker::new(file_schema, true).with_cast_collection();
let _ = expr.visit(&mut checker);
let columns = checker.into_sorted_columns();
all_root_indices.extend_from_slice(&columns.required_columns);
all_struct_accesses.extend(columns.struct_field_accesses);
all_cast_accesses.extend(columns.cast_accesses);
}
all_root_indices.sort_unstable();
all_root_indices.dedup();
all_cast_accesses.retain(|c| all_root_indices.binary_search(&c.root_index).is_err());
if !all_cast_accesses.is_empty() {
return build_read_plan_with_cast_clipping(
file_schema,
schema_descr,
&all_root_indices,
&all_struct_accesses,
&all_cast_accesses,
);
}
if all_struct_accesses.is_empty() {
return root_level_plan(&all_root_indices, file_schema, schema_descr);
}
let (read_plan, _leaf_indices) = assemble_read_plan(
&all_root_indices,
&all_struct_accesses,
file_schema,
schema_descr,
);
read_plan
}
fn build_read_plan_with_cast_clipping(
file_schema: &Schema,
schema_descr: &SchemaDescriptor,
whole_root_indices: &[usize],
struct_accesses: &[StructFieldAccess],
cast_accesses: &[CastColumnAccess],
) -> ParquetReadPlan {
let whole_roots: BTreeSet<usize> = whole_root_indices.iter().copied().collect();
let struct_access_roots: BTreeSet<usize> =
struct_accesses.iter().map(|a| a.root_index).collect();
let leaves_by_root = leaves_grouped_by_root(schema_descr);
let mut clipped_by_root: BTreeMap<usize, (Vec<usize>, DataType)> = BTreeMap::new();
let mut fallback_roots: BTreeSet<usize> = BTreeSet::new();
let mut clipped_target_by_root: BTreeMap<usize, &DataType> = BTreeMap::new();
for access in cast_accesses {
let root = access.root_index;
if whole_roots.contains(&root) || fallback_roots.contains(&root) {
continue;
}
if let Some(previous) = clipped_target_by_root.get(&root) {
if **previous != access.target_type {
clipped_by_root.remove(&root);
clipped_target_by_root.remove(&root);
fallback_roots.insert(root);
}
continue;
}
if struct_access_roots.contains(&root) {
fallback_roots.insert(root);
continue;
}
let physical_type = file_schema.field(root).data_type();
let root_leaves = leaves_by_root.get(&root).map_or(&[][..], Vec::as_slice);
if root_leaves.len() != count_leaves(physical_type) {
fallback_roots.insert(root);
continue;
}
match clip_for_cast(physical_type, &access.target_type) {
Some((kept_offsets, pruned_type)) => {
let start = root_leaves[0];
let absolute = kept_offsets.into_iter().map(|o| start + o).collect();
clipped_by_root.insert(root, (absolute, pruned_type));
clipped_target_by_root.insert(root, &access.target_type);
}
None => {
fallback_roots.insert(root);
}
}
}
let get_field_accesses: Vec<StructFieldAccess> = struct_accesses
.iter()
.filter(|a| {
debug_assert!(!clipped_by_root.contains_key(&a.root_index));
!whole_roots.contains(&a.root_index)
&& !fallback_roots.contains(&a.root_index)
})
.cloned()
.collect();
let mut leaf_indices: Vec<usize> = Vec::new();
let mut fields: BTreeMap<usize, Arc<Field>> = BTreeMap::new();
for root in whole_roots.iter().chain(fallback_roots.iter()) {
if let Some(leaves) = leaves_by_root.get(root) {
leaf_indices.extend(leaves.iter().copied());
}
fields.insert(*root, Arc::new(file_schema.field(*root).clone()));
}
for (&root, (kept, pruned_type)) in &clipped_by_root {
leaf_indices.extend(kept.iter().copied());
fields.insert(
root,
field_with_type(file_schema.field(root), pruned_type.clone()),
);
}
if !get_field_accesses.is_empty() {
let get_field_tree = StructAccessTree::from_accesses(&get_field_accesses);
leaf_indices.extend(resolve_struct_field_leaves(&get_field_tree, schema_descr));
let get_field_schema = build_filter_schema(file_schema, &[], &get_field_tree);
let get_field_roots: BTreeSet<usize> =
get_field_accesses.iter().map(|a| a.root_index).collect();
debug_assert_eq!(get_field_roots.len(), get_field_schema.fields().len());
for (root, field) in get_field_roots.iter().zip(get_field_schema.fields()) {
fields.insert(*root, Arc::clone(field));
}
}
leaf_indices.sort_unstable();
leaf_indices.dedup();
ParquetReadPlan {
projection_mask: ProjectionMask::leaves(
schema_descr,
leaf_indices.iter().copied(),
),
projected_schema: Arc::new(Schema::new_with_metadata(
fields.into_values().collect::<Vec<_>>(),
file_schema.metadata().clone(),
)),
}
}
fn leaves_grouped_by_root(
schema_descr: &SchemaDescriptor,
) -> BTreeMap<usize, Vec<usize>> {
let mut by_root: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
for leaf_idx in 0..schema_descr.num_columns() {
by_root
.entry(schema_descr.get_column_root_idx(leaf_idx))
.or_default()
.push(leaf_idx);
}
by_root
}
pub(crate) fn assemble_read_plan(
root_indices: &[usize],
struct_field_accesses: &[StructFieldAccess],
file_schema: &Schema,
schema_descr: &SchemaDescriptor,
) -> (ParquetReadPlan, Vec<usize>) {
let access_tree = StructAccessTree::from_accesses(struct_field_accesses);
let mut leaf_indices =
leaf_indices_for_roots(root_indices.iter().copied(), schema_descr);
leaf_indices
.extend_from_slice(&resolve_struct_field_leaves(&access_tree, schema_descr));
leaf_indices.sort_unstable();
leaf_indices.dedup();
let projection_mask =
ProjectionMask::leaves(schema_descr, leaf_indices.iter().copied());
let projected_schema = build_filter_schema(file_schema, root_indices, &access_tree);
(
ParquetReadPlan {
projection_mask,
projected_schema,
},
leaf_indices,
)
}
fn root_level_plan(
root_indices: &[usize],
file_schema: &Schema,
schema_descr: &SchemaDescriptor,
) -> ParquetReadPlan {
let projection_mask =
ProjectionMask::roots(schema_descr, root_indices.iter().copied());
let projected_schema = Arc::new(
file_schema
.project(root_indices)
.expect("valid column indices"),
);
ParquetReadPlan {
projection_mask,
projected_schema,
}
}
fn leaf_indices_for_roots<I>(
root_indices: I,
schema_descr: &SchemaDescriptor,
) -> Vec<usize>
where
I: IntoIterator<Item = usize>,
{
let root_set: BTreeSet<_> = root_indices.into_iter().collect();
(0..schema_descr.num_columns())
.filter(|leaf_idx| {
root_set.contains(&schema_descr.get_column_root_idx(*leaf_idx))
})
.collect()
}
fn resolve_struct_field_leaves(
access_tree: &StructAccessTree<'_>,
schema_descr: &SchemaDescriptor,
) -> Vec<usize> {
let mut leaf_indices = Vec::new();
for leaf_idx in 0..schema_descr.num_columns() {
let root_idx = schema_descr.get_column_root_idx(leaf_idx);
let Some(root_node) = access_tree.roots.get(&root_idx) else {
continue;
};
let col = schema_descr.column(leaf_idx);
let Some((_root_name, rest)) = col.path().parts().split_first() else {
continue;
};
if leaf_under_tree(root_node, rest) {
leaf_indices.push(leaf_idx);
}
}
leaf_indices
}
fn leaf_under_tree(mut node: &StructAccessNode<'_>, path: &[String]) -> bool {
for component in path {
if node.selected_here {
return true;
}
let Some(child) = node.children.get(component.as_str()) else {
return false;
};
node = child;
}
node.selected_here
}
fn build_filter_schema(
file_schema: &Schema,
regular_indices: &[usize],
access_tree: &StructAccessTree<'_>,
) -> SchemaRef {
let regular_set: BTreeSet<usize> = regular_indices.iter().copied().collect();
let all_indices = regular_indices
.iter()
.copied()
.chain(access_tree.roots.keys().copied())
.collect::<BTreeSet<_>>();
let fields = all_indices
.iter()
.map(|&idx| {
let field = file_schema.field(idx);
if regular_set.contains(&idx) {
return Arc::new(field.clone());
}
let Some(node) = access_tree.root(idx) else {
return Arc::new(field.clone());
};
let pruned_data_type = prune_struct_type(field.data_type(), node);
Arc::new(Field::new(
field.name(),
pruned_data_type,
field.is_nullable(),
))
})
.collect::<Vec<_>>();
Arc::new(Schema::new_with_metadata(
fields,
file_schema.metadata().clone(),
))
}
fn prune_struct_type(dt: &DataType, node: &StructAccessNode<'_>) -> DataType {
if node.selected_here {
return dt.clone();
}
let DataType::Struct(fields) = dt else {
return dt.clone();
};
let pruned_fields = fields
.iter()
.filter_map(|f| {
let child = node.children.get(f.name().as_str())?;
let out = if child.selected_here {
Arc::clone(f)
} else {
let pruned = prune_struct_type(f.data_type(), child);
Arc::new(Field::new(f.name(), pruned, f.is_nullable()))
};
Some(out)
})
.collect::<Vec<_>>();
DataType::Struct(pruned_fields.into())
}
#[cfg(test)]
mod test {
use super::*;
use Column as PhysicalColumn;
use arrow::array::{Int32Array, RecordBatch, StringArray, StructArray};
use arrow::datatypes::Fields;
use datafusion_common::ScalarValue;
use datafusion_expr::{Expr, col};
use datafusion_functions::core::get_field;
use datafusion_physical_expr::planner::logical2physical;
use parquet::arrow::ArrowWriter;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
use parquet::file::metadata::ParquetMetaData;
use tempfile::NamedTempFile;
#[test]
fn projection_read_plan_preserves_full_struct() {
let struct_fields: Fields = vec![
Arc::new(Field::new("value", DataType::Int32, false)),
Arc::new(Field::new("label", DataType::Utf8, false)),
]
.into();
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("s", DataType::Struct(struct_fields.clone()), false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StructArray::new(
struct_fields,
vec![
Arc::new(Int32Array::from(vec![10, 20, 30])) as _,
Arc::new(StringArray::from(vec!["a", "b", "c"])) as _,
],
None,
)),
],
)
.unwrap();
let file = NamedTempFile::new().expect("temp file");
let mut writer =
ArrowWriter::try_new(file.reopen().unwrap(), Arc::clone(&schema), None)
.expect("writer");
writer.write(&batch).expect("write batch");
writer.close().expect("close writer");
let reader_file = file.reopen().expect("reopen file");
let builder = ParquetRecordBatchReaderBuilder::try_new(reader_file)
.expect("reader builder");
let metadata = builder.metadata().clone();
let file_schema = builder.schema().clone();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![
Arc::new(PhysicalColumn::new("id", 0)),
Arc::new(PhysicalColumn::new("s", 1)),
logical2physical(
&get_field().call(vec![
col("s"),
Expr::Literal(ScalarValue::Utf8(Some("value".to_string())), None),
]),
&file_schema,
),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(
s_field.data_type(),
&DataType::Struct(
vec![
Arc::new(Field::new("value", DataType::Int32, false)),
Arc::new(Field::new("label", DataType::Utf8, false)),
]
.into()
),
);
let expected_mask = ProjectionMask::leaves(schema_descr, [0, 1, 2]);
assert_eq!(read_plan.projection_mask, expected_mask,);
}
fn write_id_struct_file() -> (SchemaRef, Arc<ParquetMetaData>) {
let struct_fields: Fields = vec![
Arc::new(Field::new("value", DataType::Int32, false)),
Arc::new(Field::new("label", DataType::Utf8, false)),
Arc::new(Field::new("pad", DataType::Utf8, false)),
]
.into();
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("s", DataType::Struct(struct_fields.clone()), false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StructArray::new(
struct_fields,
vec![
Arc::new(Int32Array::from(vec![10, 20, 30])) as _,
Arc::new(StringArray::from(vec!["a", "b", "c"])) as _,
Arc::new(StringArray::from(vec!["p0", "p1", "p2"])) as _,
],
None,
)),
],
)
.unwrap();
let file = NamedTempFile::new().expect("temp file");
let mut writer =
ArrowWriter::try_new(file.reopen().unwrap(), Arc::clone(&schema), None)
.expect("writer");
writer.write(&batch).expect("write batch");
writer.close().expect("close writer");
let builder = ParquetRecordBatchReaderBuilder::try_new(file.reopen().unwrap())
.expect("reader builder");
(builder.schema().clone(), builder.metadata().clone())
}
fn write_two_struct_file() -> (SchemaRef, Arc<ParquetMetaData>) {
let group = |first: &str, second: &str| -> Fields {
vec![
Arc::new(Field::new(first, DataType::Int32, false)),
Arc::new(Field::new(second, DataType::Utf8, false)),
]
.into()
};
let (a_fields, b_fields) = (group("p", "q"), group("m", "n"));
let schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Struct(a_fields.clone()), false),
Field::new("b", DataType::Struct(b_fields.clone()), false),
]));
let values = |fields: Fields, ints: [i32; 2], strs: [&str; 2]| {
Arc::new(StructArray::new(
fields,
vec![
Arc::new(Int32Array::from(ints.to_vec())) as _,
Arc::new(StringArray::from(strs.to_vec())) as _,
],
None,
)) as _
};
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
values(a_fields, [1, 2], ["a0", "a1"]),
values(b_fields, [3, 4], ["b0", "b1"]),
],
)
.unwrap();
let file = NamedTempFile::new().expect("temp file");
let mut writer =
ArrowWriter::try_new(file.reopen().unwrap(), Arc::clone(&schema), None)
.expect("writer");
writer.write(&batch).expect("write batch");
writer.close().expect("close writer");
let builder = ParquetRecordBatchReaderBuilder::try_new(file.reopen().unwrap())
.expect("reader builder");
(builder.schema().clone(), builder.metadata().clone())
}
fn cast_to_struct(
name: &str,
index: usize,
fields: Vec<(&str, DataType)>,
) -> Arc<dyn PhysicalExpr> {
let target = DataType::Struct(
fields
.into_iter()
.map(|(n, dt)| Arc::new(Field::new(n, dt, true)))
.collect::<Vec<_>>()
.into(),
);
Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new(name, index)),
target,
None,
))
}
fn get_field_of(
file_schema: &Schema,
name: &str,
field: &str,
) -> Arc<dyn PhysicalExpr> {
logical2physical(
&get_field().call(vec![
col(name),
Expr::Literal(ScalarValue::Utf8(Some(field.to_string())), None),
]),
file_schema,
)
}
#[test]
fn build_projection_read_plan_clips_cast_to_a_non_leading_field() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs = vec![cast_to_struct("s", 1, vec![("label", DataType::Utf8)])];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [2])
);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(
s_field.data_type(),
&DataType::Struct(
vec![Arc::new(Field::new("label", DataType::Utf8, false))].into()
),
);
}
#[test]
fn build_projection_read_plan_clips_cast_beside_get_field_on_another_root() {
let (file_schema, metadata) = write_two_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs = vec![
cast_to_struct("a", 0, vec![("p", DataType::Int32)]),
get_field_of(&file_schema, "b", "n"),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [0, 3])
);
let field_types = read_plan
.projected_schema
.fields()
.iter()
.map(|f| (f.name().clone(), f.data_type().clone()))
.collect::<Vec<_>>();
assert_eq!(
field_types,
vec![
(
"a".to_string(),
DataType::Struct(
vec![Arc::new(Field::new("p", DataType::Int32, false))].into()
)
),
(
"b".to_string(),
DataType::Struct(
vec![Arc::new(Field::new("n", DataType::Utf8, false))].into()
)
),
]
);
}
#[test]
fn build_projection_read_plan_keeps_full_read_after_a_third_cast() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs = vec![
cast_to_struct("s", 1, vec![("value", DataType::Int32)]),
cast_to_struct("s", 1, vec![("label", DataType::Utf8)]),
cast_to_struct("s", 1, vec![("value", DataType::Int32)]),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [1, 2, 3])
);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(s_field.data_type(), file_schema.field(1).data_type());
}
#[test]
fn build_projection_read_plan_whole_column_beats_get_field_beside_a_clip() {
let (file_schema, metadata) = write_two_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![
Arc::new(PhysicalColumn::new("a", 0)),
get_field_of(&file_schema, "a", "p"),
cast_to_struct("b", 1, vec![("m", DataType::Int32)]),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [0, 1, 2])
);
let a_field = read_plan.projected_schema.field_with_name("a").unwrap();
assert_eq!(
a_field.data_type(),
file_schema.field(0).data_type(),
"the whole-column reference must keep `a`'s full type"
);
}
#[test]
fn build_projection_read_plan_resolves_stale_column_indices_by_name() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs = vec![cast_to_struct("s", 0, vec![("value", DataType::Int32)])];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [1]),
"the cast must resolve to `s`, not to whatever sits at index 0"
);
}
#[test]
fn build_projection_read_plan_clips_cast_over_struct() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let narrow = DataType::Struct(
vec![Arc::new(Field::new("value", DataType::Int32, true))].into(),
);
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![
Arc::new(PhysicalColumn::new("id", 0)),
Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new("s", 1)),
narrow.clone(),
None,
)),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
let expected_mask = ProjectionMask::leaves(schema_descr, [0, 1]);
assert_eq!(read_plan.projection_mask, expected_mask);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(
s_field.data_type(),
&DataType::Struct(
vec![Arc::new(Field::new("value", DataType::Int32, false))].into()
),
);
}
#[test]
fn build_projection_read_plan_clips_repeated_identical_casts() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let narrow = DataType::Struct(
vec![Arc::new(Field::new("value", DataType::Int32, true))].into(),
);
let cast = || -> Arc<dyn PhysicalExpr> {
Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new("s", 1)),
narrow.clone(),
None,
))
};
let read_plan =
build_projection_read_plan(vec![cast(), cast()], &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [1])
);
}
#[test]
fn build_projection_read_plan_falls_back_on_conflicting_cast_targets() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let narrow = |name: &str, dt: DataType| -> Arc<dyn PhysicalExpr> {
Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new("s", 1)),
DataType::Struct(vec![Arc::new(Field::new(name, dt, true))].into()),
None,
))
};
let exprs = vec![
narrow("value", DataType::Int32),
narrow("label", DataType::Utf8),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::leaves(schema_descr, [1, 2, 3]),
"every leaf of `s` must be read so both casts see their fields"
);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(s_field.data_type(), file_schema.field(1).data_type());
}
#[test]
fn build_projection_read_plan_ignores_unprojected_struct_columns() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new("id", 0)),
DataType::Int64,
None,
))];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
assert_eq!(
read_plan.projection_mask,
ProjectionMask::roots(schema_descr, [0])
);
assert_eq!(read_plan.projected_schema.fields().len(), 1);
}
#[test]
fn build_projection_read_plan_falls_back_when_cast_and_get_field_share_a_root() {
let (file_schema, metadata) = write_id_struct_file();
let schema_descr = metadata.file_metadata().schema_descr();
let narrow = DataType::Struct(
vec![Arc::new(Field::new("value", DataType::Int32, true))].into(),
);
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![
Arc::new(CastExpr::new(
Arc::new(PhysicalColumn::new("s", 1)),
narrow,
None,
)),
logical2physical(
&get_field().call(vec![
col("s"),
Expr::Literal(ScalarValue::Utf8(Some("label".to_string())), None),
]),
&file_schema,
),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
let expected_mask = ProjectionMask::leaves(schema_descr, [1, 2, 3]);
assert_eq!(read_plan.projection_mask, expected_mask);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(
s_field.data_type(),
&DataType::Struct(
vec![
Arc::new(Field::new("value", DataType::Int32, false)),
Arc::new(Field::new("label", DataType::Utf8, false)),
Arc::new(Field::new("pad", DataType::Utf8, false)),
]
.into()
),
);
}
fn access(root: usize, path: &[&str]) -> StructFieldAccess {
StructFieldAccess {
root_index: root,
field_path: path.iter().map(|&s| s.to_string()).collect(),
}
}
#[test]
fn struct_access_tree_from_empty_input_has_no_roots() {
let tree = StructAccessTree::from_accesses(&[]);
assert!(tree.roots.is_empty());
}
#[test]
fn struct_access_tree_groups_paths_by_root() {
let accesses = [access(0, &["a"]), access(2, &["x"]), access(2, &["y"])];
let tree = StructAccessTree::from_accesses(&accesses);
assert_eq!(tree.roots.keys().copied().collect::<Vec<_>>(), vec![0, 2]);
let root0 = tree.root(0).unwrap();
assert!(root0.children.contains_key("a"));
assert!(root0.children["a"].selected_here);
let root2 = tree.root(2).unwrap();
assert_eq!(
root2.children.keys().copied().collect::<Vec<_>>(),
vec!["x", "y"],
);
}
#[test]
fn struct_access_tree_shared_prefix_collapses_into_one_node() {
let accesses = [access(0, &["outer", "a"]), access(0, &["outer", "b"])];
let tree = StructAccessTree::from_accesses(&accesses);
let root = tree.root(0).unwrap();
assert!(!root.selected_here);
let outer = &root.children["outer"];
assert!(!outer.selected_here);
assert_eq!(
outer.children.keys().copied().collect::<Vec<_>>(),
vec!["a", "b"],
);
assert!(outer.children["a"].selected_here);
assert!(outer.children["b"].selected_here);
}
#[test]
fn struct_access_tree_records_both_shallow_and_deep_selection() {
let accesses = [access(0, &["outer"]), access(0, &["outer", "a"])];
let tree = StructAccessTree::from_accesses(&accesses);
let outer = &tree.root(0).unwrap().children["outer"];
assert!(outer.selected_here);
assert!(outer.children["a"].selected_here);
}
#[test]
fn prune_struct_type_returns_full_type_when_node_is_selected_here() {
let node = StructAccessNode {
selected_here: true,
..Default::default()
};
let s_type = DataType::Struct(
vec![
Arc::new(Field::new("outer", DataType::Int32, false)),
Arc::new(Field::new("other", DataType::Int32, false)),
]
.into(),
);
let pruned = prune_struct_type(&s_type, &node);
assert_eq!(
pruned, s_type,
"selected_here on the input node must preserve the full type"
);
}
#[test]
fn prune_struct_type_shallow_selection_subsumes_deeper_children() {
let accesses = [access(0, &["outer"]), access(0, &["outer", "a"])];
let tree = StructAccessTree::from_accesses(&accesses);
let outer_type = DataType::Struct(
vec![
Arc::new(Field::new("a", DataType::Int32, false)),
Arc::new(Field::new("b", DataType::Int32, false)),
]
.into(),
);
let outer_node = &tree.root(0).unwrap().children["outer"];
let pruned = prune_struct_type(&outer_type, outer_node);
assert_eq!(
pruned, outer_type,
"shallow selected_here must preserve the whole subtree, \
not narrow to the deeper child"
);
}
#[test]
fn projection_whole_root_plus_nested_access_keeps_full_struct() {
let outer_fields: Fields = vec![
Arc::new(Field::new("a", DataType::Int32, false)),
Arc::new(Field::new("b", DataType::Int32, false)),
]
.into();
let s_fields: Fields = vec![Arc::new(Field::new(
"outer",
DataType::Struct(outer_fields.clone()),
false,
))]
.into();
let schema = Arc::new(Schema::new(vec![Field::new(
"s",
DataType::Struct(s_fields.clone()),
false,
)]));
let outer_arr = StructArray::new(
outer_fields.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2])) as _,
Arc::new(Int32Array::from(vec![3, 4])) as _,
],
None,
);
let s_arr =
StructArray::new(s_fields.clone(), vec![Arc::new(outer_arr) as _], None);
let batch =
RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(s_arr)]).unwrap();
let file = NamedTempFile::new().expect("temp file");
let mut writer =
ArrowWriter::try_new(file.reopen().unwrap(), Arc::clone(&schema), None)
.expect("writer");
writer.write(&batch).expect("write batch");
writer.close().expect("close writer");
let reader_file = file.reopen().expect("reopen file");
let builder = ParquetRecordBatchReaderBuilder::try_new(reader_file)
.expect("reader builder");
let metadata = builder.metadata().clone();
let file_schema = builder.schema().clone();
let schema_descr = metadata.file_metadata().schema_descr();
let exprs: Vec<Arc<dyn PhysicalExpr>> = vec![
Arc::new(PhysicalColumn::new("s", 0)),
logical2physical(
&get_field().call(vec![
col("s"),
Expr::Literal(ScalarValue::Utf8(Some("outer".to_string())), None),
Expr::Literal(ScalarValue::Utf8(Some("a".to_string())), None),
]),
&file_schema,
),
];
let read_plan = build_projection_read_plan(exprs, &file_schema, schema_descr);
let s_field = read_plan.projected_schema.field_with_name("s").unwrap();
assert_eq!(
s_field.data_type(),
&DataType::Struct(s_fields),
"whole-root reference must preserve the full nested struct type \
even when a nested access is also recorded"
);
let expected_mask = ProjectionMask::leaves(schema_descr, [0, 1]);
assert_eq!(
read_plan.projection_mask, expected_mask,
"whole-root reference must select every leaf under the root"
);
}
}