use std::collections::HashSet;
use std::sync::Arc;
use datafusion::common::config::ConfigOptions;
use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion};
use datafusion::error::{DataFusionError, Result as DataFusionResult};
use datafusion::execution::{SendableRecordBatchStream, TaskContext};
use datafusion::physical_expr::expressions::Column;
use datafusion::physical_expr::{PhysicalExpr, PhysicalSortExpr};
use datafusion::physical_plan::filter_pushdown::{
ChildFilterDescription, FilterDescription, FilterPushdownPhase,
};
use datafusion::physical_plan::sort_pushdown::SortOrderPushdownResult;
use datafusion::physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, Statistics,
};
#[derive(Debug)]
pub struct NanPruningBarrierExec {
input: Arc<dyn ExecutionPlan>,
unsafe_columns: Arc<HashSet<String>>,
properties: Arc<PlanProperties>,
}
impl NanPruningBarrierExec {
pub fn new(input: Arc<dyn ExecutionPlan>, unsafe_columns: Arc<HashSet<String>>) -> Self {
let properties = Arc::clone(input.properties());
Self {
input,
unsafe_columns,
properties,
}
}
}
impl DisplayAs for NanPruningBarrierExec {
fn fmt_as(&self, _t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
let mut columns: Vec<&str> = self.unsafe_columns.iter().map(String::as_str).collect();
columns.sort_unstable();
write!(
f,
"NanPruningBarrierExec: unsafe_columns=[{}]",
columns.join(", ")
)
}
}
impl ExecutionPlan for NanPruningBarrierExec {
fn name(&self) -> &str {
"NanPruningBarrierExec"
}
fn properties(&self) -> &Arc<PlanProperties> {
&self.properties
}
fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
vec![&self.input]
}
fn maintains_input_order(&self) -> Vec<bool> {
vec![true]
}
fn supports_limit_pushdown(&self) -> bool {
true
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
if children.len() != 1 {
return Err(DataFusionError::Internal(
"NanPruningBarrierExec expects exactly one child".into(),
));
}
Ok(Arc::new(NanPruningBarrierExec::new(
Arc::clone(&children[0]),
Arc::clone(&self.unsafe_columns),
)))
}
fn gather_filters_for_pushdown(
&self,
_phase: FilterPushdownPhase,
parent_filters: Vec<Arc<dyn PhysicalExpr>>,
_config: &ConfigOptions,
) -> DataFusionResult<FilterDescription> {
let allowed: HashSet<usize> = self
.input
.schema()
.fields()
.iter()
.enumerate()
.filter(|(_, field)| !self.unsafe_columns.contains(field.name()))
.map(|(index, _)| index)
.collect();
let child = ChildFilterDescription::from_child_with_allowed_indices(
&parent_filters,
allowed,
&self.input,
)?;
Ok(FilterDescription::new().with_child(child))
}
fn try_pushdown_sort(
&self,
order: &[PhysicalSortExpr],
) -> DataFusionResult<SortOrderPushdownResult<Arc<dyn ExecutionPlan>>> {
let references_unsafe = order.iter().any(|sort| {
let mut found = false;
sort.expr
.apply(|expr| {
if let Some(column) = expr.downcast_ref::<Column>()
&& self.unsafe_columns.contains(column.name())
{
found = true;
return Ok(TreeNodeRecursion::Stop);
}
Ok(TreeNodeRecursion::Continue)
})
.expect("column scan over a sort expression cannot fail");
found
});
if references_unsafe {
return Ok(SortOrderPushdownResult::Unsupported);
}
Ok(self.input.try_pushdown_sort(order)?.map(|inner| {
Arc::new(NanPruningBarrierExec::new(
inner,
Arc::clone(&self.unsafe_columns),
)) as Arc<dyn ExecutionPlan>
}))
}
fn execute(
&self,
partition: usize,
context: Arc<TaskContext>,
) -> DataFusionResult<SendableRecordBatchStream> {
self.input.execute(partition, context)
}
fn partition_statistics(&self, partition: Option<usize>) -> DataFusionResult<Arc<Statistics>> {
self.input.partition_statistics(partition)
}
}