use std::sync::Arc;
use crate::PhysicalOptimizerRule;
use arrow::datatypes::DataType;
use datafusion_common::config::ConfigOptions;
use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode};
use datafusion_common::{Result, ScalarValue};
use datafusion_expr::Operator;
use datafusion_physical_expr::expressions::{BinaryExpr, Column, Literal};
use datafusion_physical_expr::window::StandardWindowExpr;
use datafusion_physical_expr::{LexOrdering, PhysicalSortExpr};
use datafusion_physical_plan::ExecutionPlan;
use datafusion_physical_plan::execution_plan::replace_children_if_necessary;
use datafusion_physical_plan::filter::FilterExec;
use datafusion_physical_plan::projection::ProjectionExec;
use datafusion_physical_plan::repartition::RepartitionExec;
use datafusion_physical_plan::sorts::partitioned_topk::{
PartitionedTopKExec, WindowFnKind,
};
use datafusion_physical_plan::windows::{BoundedWindowAggExec, WindowUDFExpr};
#[derive(Default, Clone, Debug)]
pub struct WindowTopN;
impl WindowTopN {
pub fn new() -> Self {
Self
}
fn try_transform(plan: &Arc<dyn ExecutionPlan>) -> Option<Arc<dyn ExecutionPlan>> {
let filter = plan.downcast_ref::<FilterExec>()?;
if filter.projection().is_some() {
return None;
}
let (col_idx, limit_n) = extract_window_limit(filter.predicate())?;
let child = filter.input();
let (window_exec, intermediates) = find_window_below(child)?;
let window_exec_typed = window_exec.downcast_ref::<BoundedWindowAggExec>()?;
let input_field_count = window_exec_typed.input().schema().fields().len();
if col_idx < input_field_count {
return None; }
let window_expr_idx = col_idx - input_field_count;
let window_exprs = window_exec_typed.window_expr();
if window_expr_idx >= window_exprs.len() {
return None;
}
let fn_kind = supported_window_fn(&window_exprs[window_expr_idx])?;
let partition_by = window_exprs[window_expr_idx].partition_by();
let partition_prefix_len = partition_by.len();
if partition_prefix_len == 0 {
return None;
}
let order_by = window_exprs[window_expr_idx].order_by();
if matches!(fn_kind, WindowFnKind::Rank) && order_by.is_empty() {
return None;
}
let expr_iterator = partition_by
.iter()
.map(|e| PhysicalSortExpr::new_default(Arc::clone(e)))
.chain(order_by.iter().cloned());
let expr = LexOrdering::new(expr_iterator)?;
let partitioned_topk = PartitionedTopKExec::try_new(
Arc::clone(window_exec_typed.input()),
expr,
partition_prefix_len,
limit_n,
fn_kind,
)
.ok()?;
let mut result =
replace_children_if_necessary(window_exec, vec![Arc::new(partitioned_topk)])
.ok()?;
for node in intermediates.into_iter().rev() {
result = replace_children_if_necessary(node, vec![result]).ok()?;
}
Some(result)
}
}
impl PhysicalOptimizerRule for WindowTopN {
fn optimize(
&self,
plan: Arc<dyn ExecutionPlan>,
config: &ConfigOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
if !config.optimizer.enable_window_topn {
return Ok(plan);
}
plan.transform_down(|node| {
Ok(
if let Some(transformed) = WindowTopN::try_transform(&node) {
Transformed::yes(transformed)
} else {
Transformed::no(node)
},
)
})
.data()
}
fn name(&self) -> &str {
"WindowTopN"
}
fn schema_check(&self) -> bool {
true
}
}
fn extract_window_limit(
predicate: &Arc<dyn datafusion_physical_expr::PhysicalExpr>,
) -> Option<(usize, usize)> {
let binary = predicate.downcast_ref::<BinaryExpr>()?;
let op = binary.op();
let left = binary.left();
let right = binary.right();
if let (Some(col), Some(lit_val)) = (
left.downcast_ref::<Column>(),
right.downcast_ref::<Literal>(),
) {
let n = scalar_to_usize(lit_val.value())?;
return match *op {
Operator::LtEq => Some((col.index(), n)),
Operator::Lt => Some((col.index(), n - 1)),
_ => None,
};
}
if let (Some(lit_val), Some(col)) = (
left.downcast_ref::<Literal>(),
right.downcast_ref::<Column>(),
) {
let n = scalar_to_usize(lit_val.value())?;
return match *op {
Operator::GtEq => Some((col.index(), n)),
Operator::Gt => Some((col.index(), n - 1)),
_ => None,
};
}
None
}
fn scalar_to_usize(value: &ScalarValue) -> Option<usize> {
if !value.data_type().is_integer() {
return None;
}
let casted = value.cast_to(&DataType::UInt64).ok()?;
match casted {
ScalarValue::UInt64(Some(v)) if v > 0 => usize::try_from(v).ok(),
_ => None,
}
}
fn supported_window_fn(
expr: &Arc<dyn datafusion_physical_expr::window::WindowExpr>,
) -> Option<WindowFnKind> {
let swe = expr.as_any().downcast_ref::<StandardWindowExpr>()?;
let swfe = swe.get_standard_func_expr();
let udf = swfe.as_any().downcast_ref::<WindowUDFExpr>()?;
match udf.fun().name() {
"row_number" => Some(WindowFnKind::RowNumber),
"rank" => Some(WindowFnKind::Rank),
_ => None,
}
}
type PlanAndIntermediates = (Arc<dyn ExecutionPlan>, Vec<Arc<dyn ExecutionPlan>>);
fn find_window_below(plan: &Arc<dyn ExecutionPlan>) -> Option<PlanAndIntermediates> {
let mut current = Arc::clone(plan);
let mut intermediates = Vec::new();
loop {
if current.downcast_ref::<BoundedWindowAggExec>().is_some() {
return Some((current, intermediates));
} else if current.downcast_ref::<ProjectionExec>().is_some()
|| current.downcast_ref::<RepartitionExec>().is_some()
{
let next = Arc::clone(current.children().first()?);
intermediates.push(current);
current = next;
} else {
return None;
}
}
}