use std::sync::Arc;
use crate::PhysicalOptimizerRule;
use datafusion_common::config::ConfigOptions;
use datafusion_common::tree_node::{
Transformed, TransformedResult, TreeNode, TreeNodeRecursion,
};
use datafusion_common::{Result, Statistics, internal_err};
use datafusion_execution::TaskContext;
use datafusion_physical_expr::Distribution;
use datafusion_physical_expr_common::sort_expr::OrderingRequirements;
use datafusion_physical_plan::execution_plan::{
Boundedness, replace_children_if_necessary,
};
use datafusion_physical_plan::projection::{
ProjectionExec, make_with_child, update_expr, update_ordering_requirement,
};
use datafusion_physical_plan::scalar_subquery::ScalarSubqueryExec;
use datafusion_physical_plan::sorts::sort::SortExec;
use datafusion_physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
use datafusion_physical_plan::{
ChildStats, ChildrenPropertiesMode, DisplayAs, DisplayFormatType, ExecutionPlan,
ExecutionPlanProperties, PlanProperties, ReplaceChildrenOptions,
SendableRecordBatchStream, StatisticsArgs,
};
#[derive(Debug)]
pub struct OutputRequirements {
mode: RuleMode,
}
impl OutputRequirements {
pub fn new_add_mode() -> Self {
Self {
mode: RuleMode::Add,
}
}
pub fn new_remove_mode() -> Self {
Self {
mode: RuleMode::Remove,
}
}
}
#[derive(Debug, Ord, PartialOrd, PartialEq, Eq, Hash)]
enum RuleMode {
Add,
Remove,
}
#[derive(Debug)]
pub struct OutputRequirementExec {
input: Arc<dyn ExecutionPlan>,
order_requirement: Option<OrderingRequirements>,
dist_requirement: Distribution,
cache: Arc<PlanProperties>,
fetch: Option<usize>,
}
impl OutputRequirementExec {
pub fn new(
input: Arc<dyn ExecutionPlan>,
requirements: Option<OrderingRequirements>,
dist_requirement: Distribution,
fetch: Option<usize>,
) -> Self {
let cache = Self::compute_properties(&input, &fetch);
Self {
input,
order_requirement: requirements,
dist_requirement,
cache: Arc::new(cache),
fetch,
}
}
pub fn input(&self) -> Arc<dyn ExecutionPlan> {
Arc::clone(&self.input)
}
fn compute_properties(
input: &Arc<dyn ExecutionPlan>,
fetch: &Option<usize>,
) -> PlanProperties {
let boundedness = if fetch.is_some() {
Boundedness::Bounded
} else {
input.boundedness()
};
PlanProperties::new(
input.equivalence_properties().clone(), input.output_partitioning().clone(), input.pipeline_behavior(), boundedness, )
}
pub fn fetch(&self) -> Option<usize> {
self.fetch
}
}
impl DisplayAs for OutputRequirementExec {
fn fmt_as(
&self,
t: DisplayFormatType,
f: &mut std::fmt::Formatter,
) -> std::fmt::Result {
match t {
DisplayFormatType::Default | DisplayFormatType::Verbose => {
let order_cols = self
.order_requirement
.as_ref()
.map(|reqs| reqs.first())
.map(|lex| {
let pairs: Vec<String> = lex
.iter()
.map(|req| {
let direction = req
.options
.as_ref()
.map(
|opt| if opt.descending { "desc" } else { "asc" },
)
.unwrap_or("unspecified");
format!("({}, {direction})", req.expr)
})
.collect();
format!("[{}]", pairs.join(", "))
})
.unwrap_or_else(|| "[]".to_string());
write!(
f,
"OutputRequirementExec: order_by={}, dist_by={}",
order_cols, self.dist_requirement
)
}
DisplayFormatType::TreeRender => {
write!(f, "")
}
}
}
}
impl ExecutionPlan for OutputRequirementExec {
fn name(&self) -> &'static str {
"OutputRequirementExec"
}
fn properties(&self) -> &Arc<PlanProperties> {
&self.cache
}
fn benefits_from_input_partitioning(&self) -> Vec<bool> {
vec![false]
}
fn required_input_distribution(&self) -> Vec<Distribution> {
self.input_distribution_requirements().into_per_child()
}
fn input_distribution_requirements(
&self,
) -> datafusion_physical_plan::InputDistributionRequirements {
datafusion_physical_plan::InputDistributionRequirements::new(vec![
self.dist_requirement.clone(),
])
}
fn maintains_input_order(&self) -> Vec<bool> {
vec![true]
}
fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
vec![&self.input]
}
fn required_input_ordering(&self) -> Vec<Option<OrderingRequirements>> {
vec![self.order_requirement.clone()]
}
fn replace_children(
self: Arc<Self>,
mut children: Vec<Arc<dyn ExecutionPlan>>,
_: ReplaceChildrenOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(Self::new(
children.remove(0), self.order_requirement.clone(),
self.dist_requirement.clone(),
self.fetch,
)))
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> Result<Arc<dyn ExecutionPlan>> {
self.replace_children(
children,
ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
)
}
fn execute(
&self,
_partition: usize,
_context: Arc<TaskContext>,
) -> Result<SendableRecordBatchStream> {
unreachable!();
}
fn child_stats_requests(&self, partition: Option<usize>) -> Vec<ChildStats> {
vec![ChildStats::At(partition)]
}
fn statistics_from_inputs(
&self,
input_stats: &[Arc<Statistics>],
_args: &StatisticsArgs,
) -> Result<Arc<Statistics>> {
Ok(Arc::clone(&input_stats[0]))
}
#[expect(
deprecated,
reason = "HashPartitioned is accepted during the KeyPartitioned migration"
)]
fn try_swapping_with_projection(
&self,
projection: &ProjectionExec,
) -> Result<Option<Arc<dyn ExecutionPlan>>> {
let proj_exprs = projection.expr();
if proj_exprs.len() >= projection.input().schema().fields().len() {
return Ok(None);
}
let mut requirements = self.required_input_ordering().swap_remove(0);
if let Some(reqs) = requirements {
let mut updated_reqs = vec![];
let (lexes, soft) = reqs.into_alternatives();
for lex in lexes.into_iter() {
let Some(updated_lex) = update_ordering_requirement(lex, proj_exprs)?
else {
return Ok(None);
};
updated_reqs.push(updated_lex);
}
requirements = OrderingRequirements::new_alternatives(updated_reqs, soft);
}
let input_distributions = self.input_distribution_requirements();
let dist_req = match input_distributions.child_distribution(0) {
Some(
Distribution::HashPartitioned(exprs)
| Distribution::KeyPartitioned(exprs),
) => {
let mut updated_exprs = vec![];
for expr in exprs {
let Some(new_expr) = update_expr(expr, projection.expr(), false)?
else {
return Ok(None);
};
updated_exprs.push(new_expr);
}
Distribution::KeyPartitioned(updated_exprs)
}
Some(dist) => dist.clone(),
None => {
return internal_err!(
"OutputRequirementExec missing input distribution requirement"
);
}
};
make_with_child(projection, &self.input()).map(|input| {
let e = OutputRequirementExec::new(input, requirements, dist_req, self.fetch);
Some(Arc::new(e) as _)
})
}
fn fetch(&self) -> Option<usize> {
self.fetch
}
fn apply_expressions(
&self,
_f: &mut dyn FnMut(
&Arc<dyn datafusion_physical_expr_common::physical_expr::PhysicalExpr>,
) -> Result<TreeNodeRecursion>,
) -> Result<TreeNodeRecursion> {
Ok(TreeNodeRecursion::Continue)
}
}
impl PhysicalOptimizerRule for OutputRequirements {
fn optimize(
&self,
plan: Arc<dyn ExecutionPlan>,
_config: &ConfigOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
match self.mode {
RuleMode::Add => require_top_ordering(plan),
RuleMode::Remove => plan
.transform_up(|plan| {
if let Some(sort_req) = plan.downcast_ref::<OutputRequirementExec>() {
Ok(Transformed::yes(sort_req.input()))
} else {
Ok(Transformed::no(plan))
}
})
.data(),
}
}
fn name(&self) -> &str {
"OutputRequirements"
}
fn schema_check(&self) -> bool {
true
}
}
fn require_top_ordering(plan: Arc<dyn ExecutionPlan>) -> Result<Arc<dyn ExecutionPlan>> {
if plan.downcast_ref::<OutputRequirementExec>().is_some() {
return Ok(plan);
}
let (new_plan, is_changed) = require_top_ordering_helper(plan)?;
if is_changed {
Ok(new_plan)
} else {
Ok(Arc::new(OutputRequirementExec::new(
new_plan,
None,
Distribution::UnspecifiedDistribution,
None,
)) as _)
}
}
fn output_requirement_child(plan: &dyn ExecutionPlan) -> Option<usize> {
if plan.children().len() == 1 {
Some(0)
} else if plan.downcast_ref::<ScalarSubqueryExec>().is_some() {
Some(0)
} else {
None
}
}
fn require_top_ordering_helper(
plan: Arc<dyn ExecutionPlan>,
) -> Result<(Arc<dyn ExecutionPlan>, bool)> {
if plan.downcast_ref::<OutputRequirementExec>().is_some() {
return Ok((plan, true));
}
if let Some(sort_exec) = plan.downcast_ref::<SortExec>() {
let req_dist = sort_exec
.input_distribution_requirements()
.into_per_child()
.swap_remove(0);
let req_ordering = sort_exec.expr();
let reqs = OrderingRequirements::from(req_ordering.clone());
let fetch = sort_exec.fetch();
Ok((
Arc::new(OutputRequirementExec::new(
plan,
Some(reqs),
req_dist,
fetch,
)) as _,
true,
))
} else if let Some(spm) = plan.downcast_ref::<SortPreservingMergeExec>() {
let reqs = OrderingRequirements::from(spm.expr().clone());
let fetch = spm.fetch();
Ok((
Arc::new(OutputRequirementExec::new(
plan,
Some(reqs),
Distribution::SinglePartition,
fetch,
)) as _,
true,
))
} else if let Some(idx) = output_requirement_child(plan.as_ref()) {
if plan.maintains_input_order()[idx]
&& plan.required_input_ordering()[idx]
.as_ref()
.is_none_or(|o| matches!(o, OrderingRequirements::Soft(_)))
{
let mut children: Vec<Arc<dyn ExecutionPlan>> =
plan.children().into_iter().map(Arc::clone).collect();
let (new_child, is_changed) =
require_top_ordering_helper(Arc::clone(&children[idx]))?;
if is_changed {
children[idx] = new_child;
return Ok((replace_children_if_necessary(plan, children)?, true));
}
}
Ok((plan, false))
} else {
Ok((plan, false))
}
}