use crate::utils::for_each_referenced_index;
use crate::{OptimizerConfig, OptimizerRule};
use datafusion_common::tree_node::{Transformed, TreeNode};
use datafusion_common::{
DFSchema, Dependency, HashSet, NullEquality, Result, ScalarValue,
};
use datafusion_expr::{
Expr, JoinType,
logical_plan::{
Aggregate, Distinct, DistinctOn, EmptyRelation, Filter, Join, Limit, LogicalPlan,
Partitioning, Projection, Repartition, Sort, SubqueryAlias,
},
};
use std::sync::Arc;
#[derive(Debug, Default, Clone)]
struct LiveColumns(HashSet<usize>);
impl LiveColumns {
fn new() -> Self {
Self(HashSet::new())
}
fn all(schema: &DFSchema) -> Self {
Self((0..schema.fields().len()).collect())
}
fn try_new<'a>(
exprs: impl IntoIterator<Item = &'a Expr>,
schema: &DFSchema,
) -> Result<Self> {
let mut live = Self::new();
live.extend_from(exprs, schema)?;
Ok(live)
}
fn extend_from<'a>(
&mut self,
exprs: impl IntoIterator<Item = &'a Expr>,
schema: &DFSchema,
) -> Result<()> {
for expr in exprs {
for_each_referenced_index(expr, schema, |idx| {
self.0.insert(idx);
})?;
}
Ok(())
}
fn insert(&mut self, idx: usize) {
self.0.insert(idx);
}
fn is_empty(&self) -> bool {
self.0.is_empty()
}
fn split_at(&self, left_len: usize) -> (Self, Self) {
let mut left = Self::new();
let mut right = Self::new();
for &idx in &self.0 {
if idx < left_len {
left.insert(idx);
} else {
right.insert(idx - left_len);
}
}
(left, right)
}
}
#[derive(Default, Debug)]
pub struct EliminateJoin;
impl EliminateJoin {
pub fn new() -> Self {
Self {}
}
}
impl OptimizerRule for EliminateJoin {
fn name(&self) -> &str {
"eliminate_join"
}
fn rewrite(
&self,
plan: LogicalPlan,
_config: &dyn OptimizerConfig,
) -> Result<Transformed<LogicalPlan>> {
let live = LiveColumns::all(plan.schema());
rewrite_subtree(plan, live, false)
}
}
fn rewrite_subtree(
plan: LogicalPlan,
live: LiveColumns,
duplicate_insensitive: bool,
) -> Result<Transformed<LogicalPlan>> {
rewrite_node(plan, live, duplicate_insensitive)?.transform_data(|plan| {
plan.map_subqueries(|subquery| {
let live = LiveColumns::all(subquery.schema());
rewrite_subtree(subquery, live, false)
})
})
}
fn rewrite_node(
plan: LogicalPlan,
live: LiveColumns,
duplicate_insensitive: bool,
) -> Result<Transformed<LogicalPlan>> {
match plan {
LogicalPlan::Join(join) => rewrite_join(join, &live, duplicate_insensitive),
LogicalPlan::Projection(Projection {
expr,
input,
schema,
..
}) => {
let child_live = LiveColumns::try_new(&expr, input.schema())?;
rewrite_single_input(input, child_live, duplicate_insensitive, |input| {
Ok(LogicalPlan::Projection(Projection::try_new_with_schema(
expr, input, schema,
)?))
})
}
LogicalPlan::Filter(Filter {
predicate, input, ..
}) => {
let mut child_live = live;
child_live.extend_from([&predicate], input.schema())?;
rewrite_single_input(input, child_live, duplicate_insensitive, |input| {
Ok(LogicalPlan::Filter(Filter::new(predicate, input)))
})
}
LogicalPlan::Aggregate(Aggregate {
input,
group_expr,
aggr_expr,
schema,
..
}) => {
let child_live = LiveColumns::try_new(
group_expr.iter().chain(&aggr_expr),
input.schema(),
)?;
let child_duplicate_insensitive =
!group_expr.is_empty() && aggr_expr.is_empty();
rewrite_single_input(
input,
child_live,
child_duplicate_insensitive,
|input| {
Ok(LogicalPlan::Aggregate(Aggregate::try_new_with_schema(
input, group_expr, aggr_expr, schema,
)?))
},
)
}
LogicalPlan::Distinct(Distinct::All(input)) => {
let child_live = LiveColumns::all(input.schema());
rewrite_single_input(input, child_live, true, |input| {
Ok(LogicalPlan::Distinct(Distinct::All(input)))
})
}
LogicalPlan::Distinct(Distinct::On(DistinctOn {
on_expr,
select_expr,
sort_expr,
input,
schema,
})) => {
let mut child_live =
LiveColumns::try_new(on_expr.iter().chain(&select_expr), input.schema())?;
if let Some(sort_expr) = &sort_expr {
child_live
.extend_from(sort_expr.iter().map(|s| &s.expr), input.schema())?;
}
rewrite_single_input(input, child_live, true, |input| {
Ok(LogicalPlan::Distinct(Distinct::On(DistinctOn {
on_expr,
select_expr,
sort_expr,
input,
schema,
})))
})
}
LogicalPlan::Sort(Sort { expr, input, fetch }) => {
let mut child_live = live;
child_live.extend_from(expr.iter().map(|s| &s.expr), input.schema())?;
let child_duplicate_insensitive = duplicate_insensitive && fetch.is_none();
rewrite_single_input(
input,
child_live,
child_duplicate_insensitive,
|input| Ok(LogicalPlan::Sort(Sort { expr, input, fetch })),
)
}
LogicalPlan::Limit(Limit { skip, fetch, input }) => {
rewrite_single_input(input, live, false, |input| {
Ok(LogicalPlan::Limit(Limit { skip, fetch, input }))
})
}
LogicalPlan::SubqueryAlias(SubqueryAlias { input, alias, .. }) => {
rewrite_single_input(input, live, duplicate_insensitive, |input| {
Ok(LogicalPlan::SubqueryAlias(SubqueryAlias::try_new(
input, alias,
)?))
})
}
LogicalPlan::Repartition(Repartition {
input,
partitioning_scheme,
}) => {
let mut child_live = live;
match &partitioning_scheme {
Partitioning::Hash(exprs, _) | Partitioning::DistributeBy(exprs) => {
child_live.extend_from(exprs, input.schema())?;
}
Partitioning::Range(range) => {
child_live.extend_from(
range.ordering().iter().map(|sort_expr| &sort_expr.expr),
input.schema(),
)?;
}
Partitioning::RoundRobinBatch(_) => {}
}
rewrite_single_input(input, child_live, duplicate_insensitive, |input| {
Ok(LogicalPlan::Repartition(Repartition {
input,
partitioning_scheme,
}))
})
}
_ => plan.map_children(|child| {
let live = LiveColumns::all(child.schema());
rewrite_subtree(child, live, false)
}),
}
}
fn rewrite_single_input<F>(
input: Arc<LogicalPlan>,
child_live: LiveColumns,
duplicate_insensitive: bool,
rebuild: F,
) -> Result<Transformed<LogicalPlan>>
where
F: FnOnce(Arc<LogicalPlan>) -> Result<LogicalPlan>,
{
rewrite_subtree(
Arc::unwrap_or_clone(input),
child_live,
duplicate_insensitive,
)?
.map_data(|input| rebuild(Arc::new(input)))
}
fn rewrite_join(
join: Join,
live: &LiveColumns,
duplicate_insensitive: bool,
) -> Result<Transformed<LogicalPlan>> {
if join.join_type == JoinType::Inner
&& join.on.is_empty()
&& matches!(
join.filter.as_ref(),
Some(Expr::Literal(ScalarValue::Boolean(Some(false)), _))
)
{
return Ok(Transformed::yes(LogicalPlan::EmptyRelation(
EmptyRelation {
produce_one_row: false,
schema: join.schema,
},
)));
}
let (visible_left, visible_right) = split_join_output_columns(&join, live);
let rewritten_join_type = match rewritten_join_type(
&join,
&visible_left,
&visible_right,
duplicate_insensitive,
) {
JoinRewrite::ReplaceWithLeft => {
let left = rewrite_subtree(
Arc::unwrap_or_clone(join.left),
visible_left,
duplicate_insensitive,
)?;
return Ok(Transformed::yes(left.data));
}
JoinRewrite::ReplaceWithRight => {
let right = rewrite_subtree(
Arc::unwrap_or_clone(join.right),
visible_right,
duplicate_insensitive,
)?;
return Ok(Transformed::yes(right.data));
}
JoinRewrite::Join(join_type) => join_type,
};
let (mut left_live, mut right_live) = match rewritten_join_type {
JoinType::LeftSemi | JoinType::LeftAnti | JoinType::LeftMark => {
(visible_left, LiveColumns::new())
}
JoinType::RightSemi | JoinType::RightAnti | JoinType::RightMark => {
(LiveColumns::new(), visible_right)
}
_ => (visible_left, visible_right),
};
add_join_condition_columns(&join, &mut left_live, &mut right_live)?;
let (left_dup_insensitive, right_dup_insensitive) =
child_duplicate_insensitivity(rewritten_join_type, duplicate_insensitive);
let left = rewrite_subtree(
Arc::unwrap_or_clone(join.left),
left_live,
left_dup_insensitive,
)?;
let right = rewrite_subtree(
Arc::unwrap_or_clone(join.right),
right_live,
right_dup_insensitive,
)?;
let changed =
left.transformed || right.transformed || rewritten_join_type != join.join_type;
let left = Arc::new(left.data);
let right = Arc::new(right.data);
if changed {
Ok(Transformed::yes(LogicalPlan::Join(Join::try_new(
left,
right,
join.on,
join.filter,
rewritten_join_type,
join.join_constraint,
join.null_equality,
join.null_aware,
)?)))
} else {
Ok(Transformed::no(LogicalPlan::Join(Join {
left,
right,
on: join.on,
filter: join.filter,
join_type: join.join_type,
join_constraint: join.join_constraint,
schema: join.schema,
null_equality: join.null_equality,
null_aware: join.null_aware,
})))
}
}
fn child_duplicate_insensitivity(
join_type: JoinType,
duplicate_insensitive: bool,
) -> (bool, bool) {
match join_type {
JoinType::Inner => (duplicate_insensitive, duplicate_insensitive),
JoinType::LeftSemi | JoinType::LeftAnti | JoinType::LeftMark => {
(duplicate_insensitive, true)
}
JoinType::RightSemi | JoinType::RightAnti | JoinType::RightMark => {
(true, duplicate_insensitive)
}
JoinType::Left | JoinType::Right | JoinType::Full => (false, false),
}
}
enum JoinRewrite {
Join(JoinType),
ReplaceWithLeft,
ReplaceWithRight,
}
fn rewritten_join_type(
join: &Join,
visible_left: &LiveColumns,
visible_right: &LiveColumns,
duplicate_insensitive: bool,
) -> JoinRewrite {
let can_remove_right = visible_right.is_empty()
&& (duplicate_insensitive
|| side_unique_on_join(
join.right.schema(),
join.on.iter().map(|(_, right)| right),
join.null_equality,
));
if join.join_type == JoinType::Left && can_remove_right {
return JoinRewrite::ReplaceWithLeft;
}
let can_remove_left = visible_left.is_empty()
&& (duplicate_insensitive
|| side_unique_on_join(
join.left.schema(),
join.on.iter().map(|(left, _)| left),
join.null_equality,
));
if join.join_type == JoinType::Right && can_remove_left {
return JoinRewrite::ReplaceWithRight;
}
if join.join_type != JoinType::Inner || join.on.is_empty() {
return JoinRewrite::Join(join.join_type);
}
if can_remove_right {
return JoinRewrite::Join(JoinType::LeftSemi);
}
if can_remove_left {
return JoinRewrite::Join(JoinType::RightSemi);
}
JoinRewrite::Join(JoinType::Inner)
}
fn add_join_condition_columns(
join: &Join,
left_live: &mut LiveColumns,
right_live: &mut LiveColumns,
) -> Result<()> {
left_live.extend_from(join.on.iter().map(|(l, _)| l), join.left.schema())?;
right_live.extend_from(join.on.iter().map(|(_, r)| r), join.right.schema())?;
if let Some(filter) = &join.filter {
left_live.extend_from([filter], join.left.schema())?;
right_live.extend_from([filter], join.right.schema())?;
}
Ok(())
}
fn split_join_output_columns(
join: &Join,
live: &LiveColumns,
) -> (LiveColumns, LiveColumns) {
let left_len = join.left.schema().fields().len();
match join.join_type {
JoinType::Inner | JoinType::Left | JoinType::Right | JoinType::Full => {
live.split_at(left_len)
}
JoinType::LeftSemi | JoinType::LeftAnti | JoinType::LeftMark => {
(live.clone(), LiveColumns::new())
}
JoinType::RightSemi | JoinType::RightAnti | JoinType::RightMark => {
(LiveColumns::new(), live.clone())
}
}
}
fn side_unique_on_join<'a>(
schema: &DFSchema,
join_exprs: impl Iterator<Item = &'a Expr>,
null_equality: NullEquality,
) -> bool {
let join_key_indices = join_exprs
.filter_map(|expr| match expr {
Expr::Alias(alias) => alias.expr.as_ref().try_as_col(),
_ => expr.try_as_col(),
})
.filter_map(|column| schema.maybe_index_of_column(column))
.collect::<Vec<usize>>();
schema.functional_dependencies().iter().any(|dependency| {
dependency.mode == Dependency::Single
&& (!dependency.nullable || null_equality == NullEquality::NullEqualsNothing)
&& dependency
.source_indices
.iter()
.all(|idx| join_key_indices.contains(idx))
})
}
#[cfg(test)]
mod tests {
use crate::OptimizerContext;
use crate::assert_optimized_plan_eq_snapshot;
use crate::eliminate_join::EliminateJoin;
use arrow::datatypes::{DataType, Field, Schema};
use datafusion_common::{
Constraint, Constraints, NullEquality, Result, ScalarValue, SplitPoint,
};
use datafusion_expr::JoinType::Inner;
use datafusion_expr::{
Expr, JoinType, Partitioning, RangePartitioning, col, exists, lit,
logical_plan::builder::{
LogicalPlanBuilder, table_scan, table_source_with_constraints,
},
out_ref_col,
};
use datafusion_functions_aggregate::expr_fn::count;
use std::sync::Arc;
macro_rules! assert_optimized_plan_equal {
(
$plan:expr,
@ $expected:literal $(,)?
) => {{
let optimizer_ctx = OptimizerContext::new().with_max_passes(1);
let rules: Vec<Arc<dyn crate::OptimizerRule + Send + Sync>> = vec![Arc::new(EliminateJoin::new())];
assert_optimized_plan_eq_snapshot!(
optimizer_ctx,
rules,
$plan,
@ $expected,
)
}};
}
#[test]
fn join_on_false() -> Result<()> {
let plan = LogicalPlanBuilder::empty(false)
.join_on(
LogicalPlanBuilder::empty(false).build()?,
Inner,
Some(lit(false)),
)?
.build()?;
assert_optimized_plan_equal!(plan, @"EmptyRelation: rows=0")
}
#[test]
fn inner_to_left_semi_when_removed_side_is_unique() -> Result<()> {
let plan = left_join_right_with_constraints(primary_key_on_id())?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn inner_to_left_semi_when_removed_side_is_unique_with_join_filter() -> Result<()> {
let right = scan("r", &test_schema(), primary_key_on_id())?;
let plan =
LogicalPlanBuilder::from(scan("l", &test_schema(), Constraints::default())?)
.join(
right,
Inner,
(vec!["l.id"], vec!["r.id"]),
Some(col("r.y").gt(col("l.x"))),
)?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
LeftSemi Join: l.id = r.id Filter: r.y > l.x
TableScan: l
TableScan: r
")
}
#[test]
fn inner_to_right_semi_when_removed_side_is_unique() -> Result<()> {
let plan = left_with_constraints_join_right(primary_key_on_id())?
.project(vec![col("r.y")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: r.y
RightSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn inner_to_left_semi_for_duplicate_insensitive_parent() -> Result<()> {
let plan = left_join_right()?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn aggregate_with_aggregates_is_not_duplicate_insensitive() -> Result<()> {
let plan = left_join_right()?
.aggregate(vec![col("l.x")], vec![count(col("l.id"))])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[count(l.id)]]
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn duplicate_insensitive_context_propagates_through_join_tree() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let middle = scan("m", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), Constraints::default())?;
let left_join_middle = LogicalPlanBuilder::from(left)
.join(middle, Inner, (vec!["l.id"], vec!["m.id"]), None)?
.build()?;
let plan = LogicalPlanBuilder::from(left_join_middle)
.join(right, Inner, (vec!["l.id"], vec!["r.id"]), None)?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
LeftSemi Join: l.id = r.id
LeftSemi Join: l.id = m.id
TableScan: l
TableScan: m
TableScan: r
")
}
#[test]
fn projection_does_not_rewrite_without_uniqueness() -> Result<()> {
let plan = left_join_right()?.project(vec![col("l.x")])?.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn required_filter_column_prevents_duplicate_insensitive_rewrite() -> Result<()> {
let plan = left_join_right()?
.filter(col("r.y").gt(lit(10_i32)))?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
Filter: r.y > Int32(10)
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn distinct_star_keeps_unreferenced_side() -> Result<()> {
let plan = left_join_right()?
.distinct()?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
Distinct:
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn distinct_drops_unreferenced_side_when_projected() -> Result<()> {
let plan = left_join_right()?
.project(vec![col("l.x")])?
.distinct()?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Distinct:
Projection: l.x
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn correlated_subquery_outer_ref_prevents_rewrite() -> Result<()> {
let subquery =
LogicalPlanBuilder::from(scan("s", &test_schema(), Constraints::default())?)
.filter(col("s.id").eq(out_ref_col(DataType::Int32, "r.y")))?
.project(vec![lit(1)])?
.build()?;
let plan = left_join_right()?
.filter(exists(Arc::new(subquery)))?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
Filter: EXISTS (<subquery>)
Subquery:
Projection: Int32(1)
Filter: s.id = outer_ref(r.y)
TableScan: s
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn inner_to_semi_inside_uncorrelated_subquery() -> Result<()> {
let subquery = left_join_right_with_constraints(primary_key_on_id())?
.project(vec![col("l.x")])?
.build()?;
let plan = LogicalPlanBuilder::from(scan(
"outer",
&test_schema(),
Constraints::default(),
)?)
.filter(exists(Arc::new(subquery)))?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Filter: EXISTS (<subquery>)
Subquery:
Projection: l.x
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
TableScan: outer
")
}
#[test]
fn inner_to_semi_inside_correlated_subquery() -> Result<()> {
let subquery = left_join_right_with_constraints(primary_key_on_id())?
.filter(col("l.x").eq(out_ref_col(DataType::Int32, "outer.id")))?
.project(vec![col("l.x")])?
.build()?;
let plan = LogicalPlanBuilder::from(scan(
"outer",
&test_schema(),
Constraints::default(),
)?)
.filter(exists(Arc::new(subquery)))?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Filter: EXISTS (<subquery>)
Subquery:
Projection: l.x
Filter: l.x = outer_ref(outer.id)
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
TableScan: outer
")
}
#[test]
fn nullable_unique_rewrites_under_null_equals_nothing() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), unique_on_x())?;
let plan = LogicalPlanBuilder::from(left)
.join(right, Inner, (vec!["l.x"], vec!["r.x"]), None)?
.project(vec![col("l.id")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.id
LeftSemi Join: l.x = r.x
TableScan: l
TableScan: r
")
}
#[test]
fn nullable_unique_does_not_rewrite_under_null_equals_null() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), unique_on_x())?;
let plan = LogicalPlanBuilder::from(left)
.join_detailed(
right,
Inner,
(vec!["l.x"], vec!["r.x"]),
None,
NullEquality::NullEqualsNull,
)?
.project(vec![col("l.id")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.id
Inner Join: l.x = r.x
TableScan: l
TableScan: r
")
}
#[test]
fn composite_unique_rewrites_when_join_covers_all_key_columns() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), composite_primary_key_on_id_x())?;
let plan = LogicalPlanBuilder::from(left)
.join(
right,
Inner,
(vec!["l.id", "l.x"], vec!["r.id", "r.x"]),
None,
)?
.project(vec![col("l.y")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.y
LeftSemi Join: l.id = r.id, l.x = r.x
TableScan: l
TableScan: r
")
}
#[test]
fn composite_unique_does_not_rewrite_when_join_misses_a_key_column() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), composite_primary_key_on_id_x())?;
let plan = LogicalPlanBuilder::from(left)
.join(right, Inner, (vec!["l.id"], vec!["r.id"]), None)?
.project(vec![col("l.y")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.y
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn top_n_sort_blocks_duplicate_insensitive_rewrite() -> Result<()> {
let plan = left_join_right()?
.sort_with_limit(vec![col("l.x").sort(true, false)], Some(5))?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
Sort: l.x ASC NULLS LAST, fetch=5
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn sort_without_fetch_preserves_duplicate_insensitive_rewrite() -> Result<()> {
let plan = left_join_right()?
.sort(vec![col("l.x").sort(true, false)])?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
Sort: l.x ASC NULLS LAST
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn limit_blocks_duplicate_insensitive_rewrite() -> Result<()> {
let plan = left_join_right()?
.limit(0, Some(5))?
.aggregate(vec![col("l.x")], Vec::<Expr>::new())?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Aggregate: groupBy=[[l.x]], aggr=[[]]
Limit: skip=0, fetch=5
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn repartition_hash_key_keeps_removed_side_live() -> Result<()> {
let plan = left_join_right_with_constraints(primary_key_on_id())?
.repartition(Partitioning::Hash(vec![col("r.y")], 4))?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
Repartition: Hash(r.y) partition_count=4
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn repartition_range_key_keeps_removed_side_live() -> Result<()> {
let plan = left_join_right_with_constraints(primary_key_on_id())?
.repartition(Partitioning::Range(RangePartitioning::try_new(
vec![col("r.y").sort(true, true)],
vec![SplitPoint::new(vec![ScalarValue::Int32(Some(10))])],
)?))?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
Repartition: Range([r.y ASC NULLS FIRST], [(10)], 2)
Inner Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn distinct_on_enables_semi_join_rewrite() -> Result<()> {
let plan = left_join_right()?
.distinct_on(vec![col("l.x")], vec![col("l.x")], None)?
.build()?;
assert_optimized_plan_equal!(plan, @r"
DistinctOn: on_expr=[[l.x]], select_expr=[[l.x]], sort_expr=[[]]
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
#[test]
fn existing_semi_join_passes_through_unchanged() -> Result<()> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), Constraints::default())?;
let plan = LogicalPlanBuilder::from(left)
.join(
right,
JoinType::LeftSemi,
(vec!["l.id"], vec!["r.id"]),
None,
)?
.project(vec![col("l.x")])?
.build()?;
assert_optimized_plan_equal!(plan, @r"
Projection: l.x
LeftSemi Join: l.id = r.id
TableScan: l
TableScan: r
")
}
fn left_join_right() -> Result<LogicalPlanBuilder> {
left_join_right_with_constraints(Constraints::default())
}
fn left_join_right_with_constraints(
right_constraints: Constraints,
) -> Result<LogicalPlanBuilder> {
let left = scan("l", &test_schema(), Constraints::default())?;
let right = scan("r", &test_schema(), right_constraints)?;
LogicalPlanBuilder::from(left).join(
right,
Inner,
(vec!["l.id"], vec!["r.id"]),
None,
)
}
fn left_with_constraints_join_right(
left_constraints: Constraints,
) -> Result<LogicalPlanBuilder> {
let left = scan("l", &test_schema(), left_constraints)?;
let right = scan("r", &test_schema(), Constraints::default())?;
LogicalPlanBuilder::from(left).join(
right,
Inner,
(vec!["l.id"], vec!["r.id"]),
None,
)
}
fn scan(
name: &str,
schema: &Schema,
constraints: Constraints,
) -> Result<datafusion_expr::logical_plan::LogicalPlan> {
if constraints.is_empty() {
table_scan(Some(name), schema, None)?.build()
} else {
LogicalPlanBuilder::scan(
name,
table_source_with_constraints(schema, constraints),
None,
)?
.build()
}
}
fn test_schema() -> Schema {
Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("x", DataType::Int32, true),
Field::new("y", DataType::Int32, true),
])
}
fn primary_key_on_id() -> Constraints {
Constraints::new_unverified(vec![Constraint::PrimaryKey(vec![0])])
}
fn unique_on_x() -> Constraints {
Constraints::new_unverified(vec![Constraint::Unique(vec![1])])
}
fn composite_primary_key_on_id_x() -> Constraints {
Constraints::new_unverified(vec![Constraint::PrimaryKey(vec![0, 1])])
}
}