use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use datafusion::common::config::ConfigOptions;
use datafusion::common::tree_node::{Transformed, TreeNode};
use datafusion::error::{DataFusionError, Result as DataFusionResult};
use datafusion::execution::{RecordBatchStream, SendableRecordBatchStream, TaskContext};
use datafusion::physical_expr::EquivalenceProperties;
use datafusion::physical_expr::PhysicalExpr;
use datafusion::physical_expr::expressions::Column;
use datafusion::physical_plan::execution_plan::Boundedness;
use datafusion::physical_plan::filter_pushdown::{
ChildFilterDescription, FilterDescription, FilterPushdownPhase,
};
use datafusion::physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, PlanProperties,
};
use futures::Stream;
#[derive(Debug)]
pub struct ColumnRenameExec {
input: Arc<dyn ExecutionPlan>,
output_schema: SchemaRef,
name_mapping: HashMap<String, String>,
reverse_mapping: Arc<HashMap<String, String>>,
properties: Arc<PlanProperties>,
}
impl ColumnRenameExec {
pub fn new(
input: Arc<dyn ExecutionPlan>,
output_schema: SchemaRef,
name_mapping: HashMap<String, String>,
) -> Self {
let eq_props = EquivalenceProperties::new(Arc::clone(&output_schema));
let properties = Arc::new(PlanProperties::new(
eq_props,
input.output_partitioning().clone(),
input.pipeline_behavior(),
Boundedness::Bounded,
));
let reverse_mapping: HashMap<String, String> = name_mapping
.iter()
.map(|(old, new)| (new.clone(), old.clone()))
.collect();
Self {
input,
output_schema,
name_mapping,
reverse_mapping: Arc::new(reverse_mapping),
properties,
}
}
pub fn is_pure_type_preserving_rename(&self) -> bool {
let input_schema = self.input.schema();
input_schema.fields().len() == self.output_schema.fields().len()
&& input_schema
.fields()
.iter()
.zip(self.output_schema.fields())
.all(|(input, output)| {
input.data_type() == output.data_type()
&& self
.reverse_mapping
.get(output.name())
.map(String::as_str)
.unwrap_or(output.name())
== input.name()
})
}
pub fn remap_filter_to_input(
&self,
filter: Arc<dyn PhysicalExpr>,
) -> DataFusionResult<Option<Arc<dyn PhysicalExpr>>> {
if !self.is_pure_type_preserving_rename() {
return Ok(None);
}
let input_schema = self.input.schema();
let output_schema = Arc::clone(&self.output_schema);
let mut valid = true;
let transformed = filter.transform_down(|expr| {
let Some(column) = expr.downcast_ref::<Column>() else {
return Ok(Transformed::no(expr));
};
let index = column.index();
if output_schema
.fields()
.get(index)
.is_none_or(|field| field.name() != column.name())
{
valid = false;
return Ok(Transformed::complete(expr));
}
Ok(Transformed::yes(Arc::new(Column::new(
input_schema.field(index).name(),
index,
))))
})?;
Ok(valid.then_some(transformed.data))
}
}
impl DisplayAs for ColumnRenameExec {
fn fmt_as(&self, _t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "ColumnRenameExec: renames={}", self.name_mapping.len())
}
}
impl ExecutionPlan for ColumnRenameExec {
fn name(&self) -> &str {
"ColumnRenameExec"
}
fn properties(&self) -> &Arc<PlanProperties> {
&self.properties
}
fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
vec![&self.input]
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
if children.len() != 1 {
return Err(DataFusionError::Internal(
"ColumnRenameExec expects exactly one child".into(),
));
}
Ok(Arc::new(ColumnRenameExec::new(
Arc::clone(&children[0]),
Arc::clone(&self.output_schema),
self.name_mapping.clone(),
)))
}
fn supports_limit_pushdown(&self) -> bool {
self.is_pure_type_preserving_rename()
}
fn gather_filters_for_pushdown(
&self,
_phase: FilterPushdownPhase,
parent_filters: Vec<Arc<dyn PhysicalExpr>>,
_config: &ConfigOptions,
) -> DataFusionResult<FilterDescription> {
if !self.is_pure_type_preserving_rename() {
return Ok(FilterDescription::new()
.with_child(ChildFilterDescription::all_unsupported(&parent_filters)));
}
let remapped = parent_filters
.iter()
.map(|filter| self.remap_filter_to_input(Arc::clone(filter)))
.collect::<DataFusionResult<Option<Vec<_>>>>()?;
let child = match remapped {
Some(remapped) => ChildFilterDescription::from_child(&remapped, &self.input)?,
None => ChildFilterDescription::all_unsupported(&parent_filters),
};
Ok(FilterDescription::new().with_child(child))
}
fn execute(
&self,
partition: usize,
context: Arc<TaskContext>,
) -> DataFusionResult<SendableRecordBatchStream> {
let input_stream = self.input.execute(partition, context)?;
Ok(Box::pin(ColumnRenameStream {
input: input_stream,
output_schema: Arc::clone(&self.output_schema),
reverse_mapping: Arc::clone(&self.reverse_mapping),
}))
}
}
struct ColumnRenameStream {
input: SendableRecordBatchStream,
output_schema: SchemaRef,
reverse_mapping: Arc<HashMap<String, String>>,
}
impl Stream for ColumnRenameStream {
type Item = DataFusionResult<RecordBatch>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match Pin::new(&mut self.input).poll_next(cx) {
Poll::Ready(Some(Ok(batch))) => {
let result: DataFusionResult<RecordBatch> =
if self.output_schema.fields().is_empty() {
use arrow::record_batch::RecordBatchOptions;
let options =
RecordBatchOptions::new().with_row_count(Some(batch.num_rows()));
RecordBatch::try_new_with_options(
Arc::clone(&self.output_schema),
vec![],
&options,
)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
} else {
let input_schema = batch.schema();
let columns: DataFusionResult<Vec<_>> = self
.output_schema
.fields()
.iter()
.map(|output_field| {
let input_name = self
.reverse_mapping
.get(output_field.name())
.map(|s| s.as_str())
.unwrap_or_else(|| output_field.name().as_str());
let idx = input_schema
.index_of(input_name)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
coerce_column(batch.column(idx), output_field.data_type())
})
.collect();
columns.and_then(|cols| {
RecordBatch::try_new(Arc::clone(&self.output_schema), cols)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
})
};
Poll::Ready(Some(result))
},
Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(e))),
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
impl RecordBatchStream for ColumnRenameStream {
fn schema(&self) -> SchemaRef {
Arc::clone(&self.output_schema)
}
}
pub(crate) fn coerce_column(
col: &arrow::array::ArrayRef,
target: &arrow::datatypes::DataType,
) -> DataFusionResult<arrow::array::ArrayRef> {
use arrow::array::{Array, make_array};
if col.data_type() == target {
return Ok(Arc::clone(col));
}
let casted = arrow::compute::cast(col, target)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
if casted.data_type() == target {
return Ok(casted);
}
let data = casted
.into_data()
.into_builder()
.data_type(target.clone())
.build()
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?;
Ok(make_array(data))
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::datatypes::{DataType, Field, Schema};
use datafusion::physical_plan::EmptyRecordBatchStream;
use datafusion::physical_plan::empty::EmptyExec;
use datafusion::physical_plan::filter_pushdown::PushedDown;
#[test]
fn test_column_rename_stream_schema() {
let input_schema = Arc::new(Schema::new(vec![Field::new(
"old_col",
DataType::Int32,
false,
)]));
let output_schema = Arc::new(Schema::new(vec![Field::new(
"new_col",
DataType::Int32,
false,
)]));
let mut reverse_mapping = HashMap::new();
reverse_mapping.insert("new_col".to_string(), "old_col".to_string());
let stream = ColumnRenameStream {
input: Box::pin(EmptyRecordBatchStream::new(input_schema)),
output_schema: Arc::clone(&output_schema),
reverse_mapping: Arc::new(reverse_mapping),
};
assert_eq!(stream.schema().field(0).name(), "new_col");
}
#[test]
fn pure_rename_remaps_filters_and_allows_limit_pushdown() {
let input_schema = Arc::new(Schema::new(vec![Field::new(
"old_col",
DataType::Int32,
false,
)]));
let output_schema = Arc::new(Schema::new(vec![Field::new(
"new_col",
DataType::Int32,
false,
)]));
let input: Arc<dyn ExecutionPlan> = Arc::new(EmptyExec::new(input_schema));
let exec = ColumnRenameExec::new(
input,
output_schema,
HashMap::from([("old_col".to_string(), "new_col".to_string())]),
);
let filter: Arc<dyn PhysicalExpr> = Arc::new(Column::new("new_col", 0));
assert!(exec.is_pure_type_preserving_rename());
assert!(exec.supports_limit_pushdown());
let remapped = exec
.remap_filter_to_input(Arc::clone(&filter))
.unwrap()
.unwrap();
let column = remapped.downcast_ref::<Column>().unwrap();
assert_eq!(column.name(), "old_col");
assert_eq!(column.index(), 0);
let description = exec
.gather_filters_for_pushdown(
FilterPushdownPhase::Pre,
vec![filter],
&ConfigOptions::new(),
)
.unwrap();
let pushed = description.parent_filters();
assert!(matches!(pushed[0][0].discriminant, PushedDown::Yes));
let column = pushed[0][0].predicate.downcast_ref::<Column>().unwrap();
assert_eq!(column.name(), "old_col");
assert_eq!(column.index(), 0);
}
#[test]
fn casts_do_not_allow_filter_or_limit_pushdown() {
let input_schema = Arc::new(Schema::new(vec![Field::new(
"old_col",
DataType::Int32,
false,
)]));
let output_schema = Arc::new(Schema::new(vec![Field::new(
"new_col",
DataType::Int64,
false,
)]));
let input: Arc<dyn ExecutionPlan> = Arc::new(EmptyExec::new(input_schema));
let exec = ColumnRenameExec::new(
input,
output_schema,
HashMap::from([("old_col".to_string(), "new_col".to_string())]),
);
let filter: Arc<dyn PhysicalExpr> = Arc::new(Column::new("new_col", 0));
assert!(!exec.is_pure_type_preserving_rename());
assert!(!exec.supports_limit_pushdown());
assert!(
exec.remap_filter_to_input(Arc::clone(&filter))
.unwrap()
.is_none()
);
let description = exec
.gather_filters_for_pushdown(
FilterPushdownPhase::Pre,
vec![filter],
&ConfigOptions::new(),
)
.unwrap();
assert!(matches!(
description.parent_filters()[0][0].discriminant,
PushedDown::No
));
}
}