From 0466146e42e83532198a0fcc2da9b72128cb4b3f Mon Sep 17 00:00:00 2001 From: kould Date: Wed, 2 Sep 2026 18:48:08 +0800 Subject: [PATCH 1/2] refactor: make scalar expressions arena-backed --- kite_sql_serde_macros/src/orm.rs | 15 +- kite_sql_serde_macros/src/projection.rs | 4 +- src/binder/aggregate.rs | 340 ++--- src/binder/analyze.rs | 2 +- src/binder/copy.rs | 2 +- src/binder/create_view.rs | 19 +- src/binder/distinct.rs | 140 +- src/binder/expr.rs | 96 +- src/binder/mod.rs | 25 +- src/binder/parser.rs | 381 ++++-- src/binder/select.rs | 592 ++++----- src/binder/update.rs | 5 +- src/binder/window.rs | 48 +- src/catalog/column.rs | 21 +- src/db.rs | 284 +--- src/execution/ddl/add_column.rs | 2 +- src/execution/ddl/create_index.rs | 2 +- src/execution/dml/analyze.rs | 4 +- src/execution/dml/copy_to_file.rs | 4 +- src/execution/dml/delete.rs | 2 +- src/execution/dml/insert.rs | 20 +- src/execution/dml/update.rs | 13 +- src/execution/dql/aggregate/hash_agg.rs | 33 +- src/execution/dql/aggregate/mod.rs | 13 +- src/execution/dql/aggregate/simple_agg.rs | 12 +- src/execution/dql/aggregate/stream_agg.rs | 55 +- .../dql/aggregate/stream_distinct.rs | 15 +- src/execution/dql/external_sort.rs | 28 +- src/execution/dql/filter.rs | 11 +- src/execution/dql/function_scan.rs | 4 +- src/execution/dql/join/hash/full_join.rs | 10 +- src/execution/dql/join/hash/inner_join.rs | 7 +- src/execution/dql/join/hash/left_join.rs | 10 +- src/execution/dql/join/hash/mod.rs | 46 +- src/execution/dql/join/hash/right_join.rs | 7 +- src/execution/dql/join/hash_join.rs | 76 +- src/execution/dql/join/nested_loop_join.rs | 107 +- src/execution/dql/mark_apply.rs | 113 +- src/execution/dql/projection.rs | 7 +- src/execution/dql/recursive_cte.rs | 50 +- src/execution/dql/sort.rs | 43 +- src/execution/dql/top_k.rs | 69 +- src/execution/dql/window.rs | 83 +- src/execution/dql/window/function.rs | 12 +- src/execution/mod.rs | 13 +- src/execution/spill/codec.rs | 9 +- src/expression/eq_col.rs | 990 ++++++++++++++ src/expression/evaluator.rs | 145 +- src/expression/function/scala.rs | 7 +- src/expression/function/table.rs | 8 +- src/expression/mod.rs | 1167 +++++++---------- src/expression/range_detacher.rs | 875 +++++++----- src/expression/simplify.rs | 411 +++--- src/expression/visitor.rs | 334 ++--- src/expression/visitor_mut.rs | 336 +++-- src/expression/window.rs | 6 +- src/function/char_length.rs | 7 +- src/function/current_date.rs | 5 +- src/function/current_timestamp.rs | 5 +- src/function/lower.rs | 7 +- src/function/numbers.rs | 8 +- src/function/octet_length.rs | 7 +- src/function/upper.rs | 7 +- src/macros/mod.rs | 8 +- src/optimizer/heuristic/optimizer.rs | 26 +- src/optimizer/rule/implementation/mod.rs | 37 +- .../rule/normalization/column_pruning.rs | 96 +- .../rule/normalization/combine_operators.rs | 124 +- .../normalization/compilation_in_advance.rs | 105 +- .../rule/normalization/elimination.rs | 97 +- .../rule/normalization/min_max_top_k.rs | 8 +- src/optimizer/rule/normalization/mod.rs | 85 +- .../rule/normalization/parameterized_index.rs | 91 +- .../rule/normalization/pushdown_predicates.rs | 246 ++-- .../rule/normalization/simplification.rs | 115 +- src/orm/ddl.rs | 39 +- src/orm/mod.rs | 734 ++++++----- src/planner/arena.rs | 115 +- src/planner/mod.rs | 79 +- src/planner/operator/aggregate.rs | 39 +- .../operator/alter_table/change_column.rs | 27 +- src/planner/operator/analyze.rs | 9 + src/planner/operator/copy_from_file.rs | 21 +- src/planner/operator/create_index.rs | 19 +- src/planner/operator/filter.rs | 22 +- src/planner/operator/join.rs | 73 +- src/planner/operator/mark_apply.rs | 27 +- src/planner/operator/mod.rs | 574 ++++---- src/planner/operator/project.rs | 19 +- src/planner/operator/recursive_cte.rs | 20 +- src/planner/operator/set_membership.rs | 17 +- src/planner/operator/sort.rs | 47 +- src/planner/operator/table_scan.rs | 30 +- src/planner/operator/top_k.rs | 21 +- src/planner/operator/union.rs | 17 +- src/planner/operator/update.rs | 25 +- src/planner/operator/visitor.rs | 114 +- src/planner/operator/visitor_mut.rs | 62 +- src/planner/operator/window.rs | 96 +- src/serdes/column.rs | 33 +- src/serdes/expression.rs | 44 + src/serdes/mod.rs | 1 + src/storage/mod.rs | 24 +- src/storage/table_codec.rs | 5 +- src/types/evaluator/cast.rs | 12 +- src/types/evaluator/tuple.rs | 10 +- src/types/index.rs | 49 +- src/types/value.rs | 3 +- tests/macros-test/src/main.rs | 38 +- tests/slt/cte.slt | 2 +- tests/slt/join.slt | 2 +- tests/slt/parameterized_subquery.slt | 173 +++ tests/slt/stream_distinct_explain.slt | 12 +- tests/slt/subquery.slt | 8 - tests/slt/update.slt | 2 +- tests/slt/where_by_index_explain.slt | 68 +- tests/slt/window.slt | 2 +- 117 files changed, 6460 insertions(+), 4491 deletions(-) create mode 100644 src/expression/eq_col.rs create mode 100644 src/serdes/expression.rs create mode 100644 tests/slt/parameterized_subquery.slt diff --git a/kite_sql_serde_macros/src/orm.rs b/kite_sql_serde_macros/src/orm.rs index b79dbfdc..ad0f004c 100644 --- a/kite_sql_serde_macros/src/orm.rs +++ b/kite_sql_serde_macros/src/orm.rs @@ -272,11 +272,11 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { let data_type = #data_type; let default = #default_tokens .map(|value| { - ::kite_sql::expression::ScalarExpression::Constant( + arena.alloc_expression(::kite_sql::expression::ScalarExpression::Constant( value .cast(&data_type) .expect("failed to cast ORM default value to column type"), - ) + )) }); let desc = ::kite_sql::catalog::column::ColumnDesc::new( data_type, @@ -449,13 +449,10 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { ] } - fn columns() -> &'static [::kite_sql::catalog::column::ColumnCatalog] { - static ORM_COLUMNS: ::std::sync::LazyLock<::std::vec::Vec<::kite_sql::catalog::column::ColumnCatalog>> = ::std::sync::LazyLock::new(|| { - vec![ - #(#orm_columns),* - ] - }); - ORM_COLUMNS.as_slice() + fn columns(arena: &mut ::kite_sql::planner::TableArena) -> ::std::vec::Vec<::kite_sql::catalog::column::ColumnCatalog> { + vec![ + #(#orm_columns),* + ] } fn indexes() -> &'static [(&'static str, &'static [&'static str], bool)] { diff --git a/kite_sql_serde_macros/src/projection.rs b/kite_sql_serde_macros/src/projection.rs index 1067fcb3..9c64c4cf 100644 --- a/kite_sql_serde_macros/src/projection.rs +++ b/kite_sql_serde_macros/src/projection.rs @@ -105,13 +105,13 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { fn bind_projection<'ctx, 'bind, 'parent, 'arena, T, A>( scope: &mut ::kite_sql::orm::ExprBindScope<'ctx, 'bind, 'parent, 'arena, T, A>, relation: &str, - ) -> ::std::result::Result<::std::vec::Vec<::kite_sql::expression::ScalarExpression>, ::kite_sql::errors::DatabaseError> + ) -> ::std::result::Result<::std::vec::Vec<::kite_sql::planner::ExprRef>, ::kite_sql::errors::DatabaseError> where T: ::kite_sql::storage::Transaction, A: AsRef<[(&'static str, ::kite_sql::types::value::DataValue)]>, { Ok(::std::vec![ - #(::kite_sql::orm::IntoOrmScalarExpression::into_orm_scalar(#projection_exprs)),* + #(#projection_exprs.into_scalar()),* ]) } } diff --git a/src/binder/aggregate.rs b/src/binder/aggregate.rs index b723fb44..254371d6 100644 --- a/src/binder/aggregate.rs +++ b/src/binder/aggregate.rs @@ -12,13 +12,11 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::collections::HashSet; - use super::{Binder, QueryBindStep}; use crate::errors::DatabaseError; use crate::expression::visitor::{walk_expr, ExprVisitor}; use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut}; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::value::DataValue; use crate::{ @@ -27,16 +25,16 @@ use crate::{ }; struct AggregateCallCollector<'a> { - agg_calls: &'a mut Vec, + agg_calls: &'a mut Vec, } -impl<'expr> ExprVisitor<'expr> for AggregateCallCollector<'_> { - fn visit(&mut self, expr: &'expr ScalarExpression) -> Result<(), DatabaseError> { - match expr { - ScalarExpression::AggCall { .. } => self.agg_calls.push(expr.clone()), - ScalarExpression::Alias { expr, .. } => self.visit(expr)?, +impl ExprVisitor> for AggregateCallCollector<'_> { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + match arena.expression(expr) { + ScalarExpression::AggCall { .. } => self.agg_calls.push(expr), + ScalarExpression::Alias { expr, .. } => self.visit(*expr, arena)?, ScalarExpression::Empty | ScalarExpression::TableFunction(_) => unreachable!(), - _ => walk_expr(self, expr)?, + _ => walk_expr(self, expr, arena)?, } Ok(()) } @@ -46,8 +44,8 @@ impl> Binder<'_, '_, T, A> pub fn bind_aggregate( &mut self, children: LogicalPlan, - agg_calls: Vec, - groupby_exprs: Vec, + agg_calls: Vec, + groupby_exprs: Vec, ) -> Result { self.context.step(QueryBindStep::Agg); Ok(AggregateOperator::build( @@ -61,46 +59,49 @@ impl> Binder<'_, '_, T, A> pub fn extract_select_aggregate( &mut self, - select_items: &mut [ScalarExpression], + select_items: &mut [ExprRef], + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { for column in select_items { - self.collect_aggregate_calls(column)?; + self.collect_aggregate_calls(*column, arena)?; } Ok(()) } pub fn extract_group_by_aggregate_exprs( &mut self, - select_list: &mut [ScalarExpression], - mut group_by_exprs: Vec, + select_list: &mut [ExprRef], + mut group_by_exprs: Vec, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.validate_groupby_illegal_column(select_list, &group_by_exprs)?; + self.validate_groupby_illegal_column(select_list, &group_by_exprs, arena)?; for expr in group_by_exprs.iter_mut() { - self.visit_group_by_expr(select_list, expr); + self.visit_group_by_expr(select_list, *expr, arena); } Ok(()) } - pub fn extract_having_orderby_aggregate_exprs( + pub fn extract_having_orderby_aggregate_exprs<'arena, I, F>( &mut self, - mut having: Option, + mut having: Option, orderby: Option, mut bind_sort_field: F, - ) -> Result<(Option, Option>), DatabaseError> + arena: &mut PlanArena<'arena>, + ) -> Result<(Option, Option>), DatabaseError> where I: IntoIterator, - F: FnMut(&mut Self, I::Item) -> Result, + F: FnMut(&mut Self, I::Item, &mut PlanArena<'arena>) -> Result, { if let Some(having) = having.as_mut() { - self.collect_aggregate_calls(having)?; + self.collect_aggregate_calls(*having, arena)?; } let mut return_orderby = None; if let Some(orderby) = orderby { let mut fields = Vec::new(); for orderby in orderby { - let field = bind_sort_field(self, orderby)?; - self.collect_aggregate_calls(&field.expr)?; + let field = bind_sort_field(self, orderby, arena)?; + self.collect_aggregate_calls(field.expr, arena)?; fields.push(field); } return_orderby = Some(fields); @@ -110,7 +111,7 @@ impl> Binder<'_, '_, T, A> pub fn bind_aggregate_output_exprs<'c>( &mut self, - exprs: impl IntoIterator, + exprs: impl IntoIterator, arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { self.bind_aggregate_output_exprs_with_outputs( @@ -123,23 +124,27 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_aggregate_output_exprs_with_outputs<'c>( &self, - agg_calls: &[ScalarExpression], - group_by_exprs: &[ScalarExpression], - exprs: impl IntoIterator, + agg_calls: &[ExprRef], + group_by_exprs: &[ExprRef], + exprs: impl IntoIterator, arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { - let mut binder = AggregateOutputBinder::new(agg_calls, group_by_exprs, arena); + let mut binder = AggregateOutputBinder::new(agg_calls, group_by_exprs); for expr in exprs { - binder.visit(expr)?; + binder.visit(expr, arena)?; } Ok(()) } - fn collect_aggregate_calls(&mut self, expr: &ScalarExpression) -> Result<(), DatabaseError> { + pub(crate) fn collect_aggregate_calls( + &mut self, + expr: ExprRef, + arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { AggregateCallCollector { agg_calls: &mut self.context.agg_calls, } - .visit(expr) + .visit(expr, arena) } /// Validate select exprs must appear in the GROUP BY clause or be used in @@ -149,16 +154,17 @@ impl> Binder<'_, '_, T, A> /// SELECT a,count(b) FROM t GROUP BY b. it's error. fn validate_groupby_illegal_column( &mut self, - select_items: &[ScalarExpression], - groupby: &[ScalarExpression], + select_items: &[ExprRef], + groupby: &[ExprRef], + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { - let mut group_raw_exprs = vec![]; + let mut unmatched_group_exprs = Vec::with_capacity(groupby.len()); for expr in groupby { - if let ScalarExpression::Alias { alias, .. } = expr { + if let ScalarExpression::Alias { alias, .. } = arena.expression(*expr) { let alias_expr = select_items.iter().find(|column| { if let ScalarExpression::Alias { alias: inner_alias, .. - } = &column + } = arena.expression(**column) { alias == inner_alias } else { @@ -167,33 +173,35 @@ impl> Binder<'_, '_, T, A> }); if let Some(inner_expr) = alias_expr { - group_raw_exprs.push(inner_expr); + unmatched_group_exprs.push(*inner_expr); } } else { - group_raw_exprs.push(expr); + unmatched_group_exprs.push(*expr); } } - let mut group_raw_set: HashSet<&ScalarExpression> = - HashSet::from_iter(group_raw_exprs.iter().copied()); for expr in select_items { - if expr.has_window_call()? { - HavingOrderByValidator::new(groupby, &self.context.agg_calls).visit(expr)?; + if expr.has_window_call(arena)? { + HavingOrderByValidator::new(groupby, &self.context.agg_calls) + .visit(*expr, arena)?; continue; } - if expr.has_agg_call()? { + if expr.has_agg_call(arena)? { continue; } - group_raw_set.remove(expr); - - if !group_raw_exprs.contains(&expr) { + let Some(position) = unmatched_group_exprs + .iter() + .position(|group_expr| expr.eq_ignore_colref_pos(*group_expr, arena)) + else { return Err(DatabaseError::AggMiss(format!( - "`{expr}` must appear in the GROUP BY clause or be used in an aggregate function" + "`{}` must appear in the GROUP BY clause or be used in an aggregate function", + expr.output_name(arena) ))); - } + }; + unmatched_group_exprs.remove(position); } - if !group_raw_set.is_empty() { + if !unmatched_group_exprs.is_empty() { return Err(DatabaseError::AggMiss( "in the GROUP BY clause the field must be in the select clause".to_string(), )); @@ -204,120 +212,132 @@ impl> Binder<'_, '_, T, A> fn visit_group_by_expr( &mut self, - select_list: &mut [ScalarExpression], - expr: &mut ScalarExpression, + select_list: &mut [ExprRef], + expr: ExprRef, + arena: &PlanArena<'_>, ) { - if let ScalarExpression::Alias { alias, .. } = expr { + if let ScalarExpression::Alias { alias, .. } = arena.expression(expr) { if let Some(i) = select_list.iter().position(|inner_expr| { if let ScalarExpression::Alias { alias: inner_alias, .. - } = &inner_expr + } = arena.expression(*inner_expr) { alias == inner_alias } else { false } }) { - self.context.group_by_exprs.push(select_list[i].clone()); + self.context.group_by_exprs.push(select_list[i]); return; } } - if let Some(i) = select_list.iter().position(|column| column == expr) { - self.context.group_by_exprs.push(select_list[i].clone()) + if let Some(i) = select_list + .iter() + .position(|column| column.eq_ignore_colref_pos(expr, arena)) + { + self.context.group_by_exprs.push(select_list[i]) } } /// Validate having or orderby clause is valid, if SQL has group by clause. - pub fn validate_having_orderby(&self, expr: &ScalarExpression) -> Result<(), DatabaseError> { + pub fn validate_having_orderby( + &self, + expr: ExprRef, + arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { if self.context.group_by_exprs.is_empty() { return Ok(()); } HavingOrderByValidator::new(&self.context.group_by_exprs, &self.context.agg_calls) - .visit(expr) + .visit(expr, arena) } } struct HavingOrderByValidator<'a> { - group_by_exprs: &'a [ScalarExpression], - agg_calls: &'a [ScalarExpression], + group_by_exprs: &'a [ExprRef], + agg_calls: &'a [ExprRef], } impl<'a> HavingOrderByValidator<'a> { - fn new(group_by_exprs: &'a [ScalarExpression], agg_calls: &'a [ScalarExpression]) -> Self { + fn new(group_by_exprs: &'a [ExprRef], agg_calls: &'a [ExprRef]) -> Self { Self { group_by_exprs, agg_calls, } } - fn agg_miss(expr: &ScalarExpression) -> DatabaseError { + fn agg_miss(expr: ExprRef, arena: &PlanArena<'_>) -> DatabaseError { DatabaseError::AggMiss(format!( - "expression '{expr}' must appear in the GROUP BY clause or be used in an aggregate function" + "expression '{}' must appear in the GROUP BY clause or be used in an aggregate function", + expr.output_name(arena) )) } } -impl<'expr> ExprVisitor<'expr> for HavingOrderByValidator<'_> { - fn visit(&mut self, expr: &'expr ScalarExpression) -> Result<(), DatabaseError> { - match expr { +impl ExprVisitor> for HavingOrderByValidator<'_> { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + let contains = |expressions: &[ExprRef]| { + expressions + .iter() + .any(|candidate| candidate.eq_ignore_colref_pos(expr, arena)) + }; + match arena.expression(expr) { ScalarExpression::AggCall { .. } => { - if self.group_by_exprs.contains(expr) || self.agg_calls.contains(expr) { + if contains(self.group_by_exprs) || contains(self.agg_calls) { Ok(()) } else { - Err(Self::agg_miss(expr)) + Err(Self::agg_miss(expr, arena)) } } ScalarExpression::ColumnRef { .. } => { - if self.group_by_exprs.contains(expr) { + if contains(self.group_by_exprs) { Ok(()) } else { - Err(Self::agg_miss(expr)) + Err(Self::agg_miss(expr, arena)) } } ScalarExpression::Alias { .. } => { - if self.group_by_exprs.contains(expr) { + if contains(self.group_by_exprs) { Ok(()) } else { - self.visit(expr.unpack_alias_ref()) + self.visit(expr.unpack_alias(arena), arena) } } ScalarExpression::Empty | ScalarExpression::TableFunction(_) => unreachable!(), - _ => walk_expr(self, expr), + _ => walk_expr(self, expr, arena), } } } -struct AggregateOutputBinder<'a, 'p> { - agg_calls: &'a [ScalarExpression], - group_by_exprs: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, +struct AggregateOutputBinder<'a> { + agg_calls: &'a [ExprRef], + group_by_exprs: &'a [ExprRef], } -impl<'a, 'p> AggregateOutputBinder<'a, 'p> { - fn new( - agg_calls: &'a [ScalarExpression], - group_by_exprs: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, - ) -> Self { +impl<'a> AggregateOutputBinder<'a> { + fn new(agg_calls: &'a [ExprRef], group_by_exprs: &'a [ExprRef]) -> Self { Self { agg_calls, group_by_exprs, - arena, } } fn output_ref( &mut self, - expr: &ScalarExpression, + expr: ExprRef, + arena: &mut PlanArena<'_>, ) -> Result, DatabaseError> { let output_count = self.agg_calls.len() + self.group_by_exprs.len(); self.agg_calls .iter() .chain(self.group_by_exprs.iter()) .position(|candidate| { - candidate == expr || candidate.unpack_alias_ref() == expr.unpack_alias_ref() + candidate.eq_ignore_colref_pos(expr, arena) + || candidate + .unpack_alias(arena) + .eq_ignore_colref_pos(expr.unpack_alias(arena), arena) }) .map(|position| { let output_expr = self @@ -331,7 +351,7 @@ impl<'a, 'p> AggregateOutputBinder<'a, 'p> { )) })?; Ok(ScalarExpression::column_expr( - output_expr.output_column_ref(self.arena), + output_expr.output_column_ref(arena), position, )) }) @@ -339,21 +359,26 @@ impl<'a, 'p> AggregateOutputBinder<'a, 'p> { } } -impl<'a> ExprVisitorMut<'a> for AggregateOutputBinder<'_, '_> { - fn visit(&mut self, expr: &'a mut ScalarExpression) -> Result<(), DatabaseError> { +impl ExprVisitorMut for AggregateOutputBinder<'_> { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { - expr: inner_expr, alias: crate::expression::AliasType::Name(_), - } = expr + .. + } = arena.expression(*expr) { - return self.visit(inner_expr); + return walk_mut_expr(self, expr, arena); } - if let Some(output_ref) = self.output_ref(expr)? { - *expr = output_ref; + if let Some(output) = self.output_ref(*expr, arena)? { + *expr = arena.alloc_expression(output); return Ok(()); } - walk_mut_expr(self, expr) + + walk_mut_expr(self, expr, arena) } } @@ -367,7 +392,7 @@ mod tests { use crate::expression::agg::AggKind; use crate::expression::visitor_mut::ExprVisitorMut; use crate::expression::{AliasType, BinaryOperator, ScalarExpression}; - use crate::planner::PlanArena; + use crate::planner::{ExprRef, PlanArena}; use crate::storage::Storage; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -380,13 +405,13 @@ mod tests { )) } - fn test_count(expr: ScalarExpression) -> ScalarExpression { - ScalarExpression::AggCall { + fn test_count(arena: &mut PlanArena, expr: ExprRef) -> ExprRef { + arena.alloc_expression(ScalarExpression::AggCall { distinct: false, kind: AggKind::Count, args: vec![expr], ty: LogicalType::Bigint, - } + }) } #[test] @@ -396,44 +421,44 @@ mod tests { let group_column = test_column(&mut arena, "c1", LogicalType::Integer); let agg_column = test_column(&mut arena, "c2", LogicalType::Integer); - let group_expr = ScalarExpression::column_expr(group_column, 0); - let agg_expr = test_count(ScalarExpression::column_expr(agg_column, 1)); + let group_expr = arena.alloc_expression(ScalarExpression::column_expr(group_column, 0)); + let agg_arg = arena.alloc_expression(ScalarExpression::column_expr(agg_column, 1)); + let agg_expr = test_count(&mut arena, agg_arg); - let agg_output = ScalarExpression::Alias { - expr: Box::new(agg_expr.clone()), + let agg_output = arena.alloc_expression(ScalarExpression::Alias { + expr: agg_expr, alias: AliasType::Name("cnt".to_string()), - }; - let group_output = ScalarExpression::Alias { - expr: Box::new(group_expr.clone()), + }); + let group_output = arena.alloc_expression(ScalarExpression::Alias { + expr: group_expr, alias: AliasType::Name("g".to_string()), - }; + }); - let mut order_by_agg = ScalarExpression::Alias { - expr: Box::new(agg_expr), + let mut order_by_agg = arena.alloc_expression(ScalarExpression::Alias { + expr: agg_expr, alias: AliasType::Name("cnt".to_string()), - }; + }); let mut order_by_group = group_expr; { let mut binder = AggregateOutputBinder::new( std::slice::from_ref(&agg_output), std::slice::from_ref(&group_output), - &mut arena, ); - binder.visit(&mut order_by_agg)?; - binder.visit(&mut order_by_group)?; + binder.visit(&mut order_by_agg, &mut arena)?; + binder.visit(&mut order_by_group, &mut arena)?; } - let expected_agg = ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( - agg_output.output_column_ref(&mut arena), - 0, - )), + let agg_column = agg_output.output_column_ref(&mut arena); + let expected_agg_inner = + arena.alloc_expression(ScalarExpression::column_expr(agg_column, 0)); + let expected_agg = arena.alloc_expression(ScalarExpression::Alias { + expr: expected_agg_inner, alias: AliasType::Name("cnt".to_string()), - }; - assert!(order_by_agg.eq_ignore_colref_pos(&expected_agg, &arena)); + }); + assert!(order_by_agg.eq_ignore_colref_pos(expected_agg, &arena)); - let expected_group = - ScalarExpression::column_expr(group_output.output_column_ref(&mut arena), 1); - assert!(order_by_group.eq_ignore_colref_pos(&expected_group, &arena)); + let group_column = group_output.output_column_ref(&mut arena); + let expected_group = arena.alloc_expression(ScalarExpression::column_expr(group_column, 1)); + assert!(order_by_group.eq_ignore_colref_pos(expected_group, &arena)); Ok(()) } @@ -443,24 +468,25 @@ mod tests { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let group_column = test_column(&mut arena, "c1", LogicalType::Integer); - let group_expr = ScalarExpression::column_expr(group_column, 0); - let group_output = ScalarExpression::Alias { - expr: Box::new(group_expr.clone()), + let group_expr = arena.alloc_expression(ScalarExpression::column_expr(group_column, 0)); + let group_output = arena.alloc_expression(ScalarExpression::Alias { + expr: group_expr, alias: AliasType::Name("g".to_string()), - }; + }); - let mut target = ScalarExpression::Alias { - expr: Box::new(ScalarExpression::Constant(1_i32.into())), - alias: AliasType::Expr(Box::new(group_expr)), - }; + let constant = arena.alloc_expression(ScalarExpression::Constant(1_i32.into())); + let mut target = arena.alloc_expression(ScalarExpression::Alias { + expr: constant, + alias: AliasType::Expr(group_expr), + }); { - let mut binder = - AggregateOutputBinder::new(&[], std::slice::from_ref(&group_output), &mut arena); - binder.visit(&mut target)?; + let mut binder = AggregateOutputBinder::new(&[], std::slice::from_ref(&group_output)); + binder.visit(&mut target, &mut arena)?; } - let expected = ScalarExpression::column_expr(group_output.output_column_ref(&mut arena), 0); - assert!(target.eq_ignore_colref_pos(&expected, &arena)); + let output_column = group_output.output_column_ref(&mut arena); + let expected = arena.alloc_expression(ScalarExpression::column_expr(output_column, 0)); + assert!(target.eq_ignore_colref_pos(expected, &arena)); Ok(()) } @@ -487,39 +513,41 @@ mod tests { let mut arena = PlanArena::new(&table_arena); let group_column = test_column(&mut arena, "c1", LogicalType::Integer); let missing_column = test_column(&mut arena, "c2", LogicalType::Integer); - let group_expr = ScalarExpression::column_expr(group_column, 0); - let missing_expr = ScalarExpression::column_expr(missing_column, 1); - binder.context.group_by_exprs.push(group_expr.clone()); + let group_expr = arena.alloc_expression(ScalarExpression::column_expr(group_column, 0)); + let missing_expr = arena.alloc_expression(ScalarExpression::column_expr(missing_column, 1)); + binder.context.group_by_exprs.push(group_expr); - binder.validate_having_orderby(&group_expr)?; - let group_alias = ScalarExpression::Alias { - expr: Box::new(group_expr.clone()), + binder.validate_having_orderby(group_expr, &arena)?; + let group_alias = arena.alloc_expression(ScalarExpression::Alias { + expr: group_expr, alias: AliasType::Name("group_alias".to_string()), - }; - binder.validate_having_orderby(&group_alias)?; + }); + binder.validate_having_orderby(group_alias, &arena)?; assert!(matches!( - binder.validate_having_orderby(&missing_expr), + binder.validate_having_orderby(missing_expr, &arena), Err(DatabaseError::AggMiss(_)) )); - let registered_agg = test_count(missing_expr.clone()); - binder.context.agg_calls.push(registered_agg.clone()); - binder.validate_having_orderby(®istered_agg)?; + let registered_agg = test_count(&mut arena, missing_expr); + binder.context.agg_calls.push(registered_agg); + binder.validate_having_orderby(registered_agg, &arena)?; + let constant = arena.alloc_expression(ScalarExpression::Constant(1_i32.into())); + let unregistered_agg = test_count(&mut arena, constant); assert!(matches!( - binder.validate_having_orderby(&test_count(ScalarExpression::Constant(1_i32.into()))), + binder.validate_having_orderby(unregistered_agg, &arena), Err(DatabaseError::AggMiss(_)) )); - let invalid_binary = ScalarExpression::Binary { + let invalid_binary = arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Eq, - left_expr: Box::new(group_expr), - right_expr: Box::new(missing_expr), + left_expr: group_expr, + right_expr: missing_expr, evaluator: None, ty: LogicalType::Boolean, - }; + }); assert!(matches!( - binder.validate_having_orderby(&invalid_binary), + binder.validate_having_orderby(invalid_binary, &arena), Err(DatabaseError::AggMiss(_)) )); diff --git a/src/binder/analyze.rs b/src/binder/analyze.rs index b5b394b0..8a4faf73 100644 --- a/src/binder/analyze.rs +++ b/src/binder/analyze.rs @@ -26,7 +26,7 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_analyze( &mut self, table_name: TableName, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result { let table = self .context diff --git a/src/binder/copy.rs b/src/binder/copy.rs index ad33ab4e..0dfff7cb 100644 --- a/src/binder/copy.rs +++ b/src/binder/copy.rs @@ -81,7 +81,7 @@ impl> Binder<'_, '_, T, A> table_name: TableName, to: bool, ext_source: ExtSource, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result { if let Some(table) = self.context.table(table_name.clone())?.cloned() { if to { diff --git a/src/binder/create_view.rs b/src/binder/create_view.rs index de0b97b7..4b12e320 100644 --- a/src/binder/create_view.rs +++ b/src/binder/create_view.rs @@ -19,7 +19,7 @@ use crate::errors::DatabaseError; use crate::expression::{AliasType, ScalarExpression}; use crate::planner::operator::create_view::CreateViewOperator; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -38,7 +38,7 @@ impl> Binder<'_, '_, T, A> mapping_schema: &[ColumnRef], arena: &mut crate::planner::PlanArena, mut column_name: impl FnMut(usize, ColumnRef, &crate::planner::PlanArena) -> String, - ) -> Vec { + ) -> Vec { let mapping_schema_len = mapping_schema.len(); let mut exprs = Vec::with_capacity(mapping_schema_len); for (i, mapping_column) in mapping_schema.iter().copied().enumerate() { @@ -54,13 +54,12 @@ impl> Binder<'_, '_, T, A> column.set_ref_table(view_name.clone(), 0, true); let output_column = arena.alloc_column(column); - exprs.push(ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr(mapping_column, i)), - alias: AliasType::Expr(Box::new(ScalarExpression::column_expr( - output_column, - i, - ))), - }); + let expr = arena.alloc_expression(ScalarExpression::column_expr(mapping_column, i)); + let alias = arena.alloc_expression(ScalarExpression::column_expr(output_column, i)); + exprs.push(arena.alloc_expression(ScalarExpression::Alias { + expr, + alias: AliasType::Expr(alias), + })); } exprs } @@ -75,7 +74,7 @@ impl> Binder<'_, '_, T, A> ))); } - let exprs: Vec = if column_names.is_empty() { + let exprs = if column_names.is_empty() { projection_exprs(&view_name, mapping_schema, arena, |i, column, arena| { output_aliases .get(i) diff --git a/src/binder/distinct.rs b/src/binder/distinct.rs index d54286b9..3ce8927a 100644 --- a/src/binder/distinct.rs +++ b/src/binder/distinct.rs @@ -18,7 +18,7 @@ use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut}; use crate::expression::ScalarExpression; use crate::planner::operator::aggregate::AggregateOperator; use crate::planner::operator::sort::SortField; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -26,7 +26,7 @@ impl> Binder<'_, '_, T, A> pub fn bind_distinct( &mut self, children: LogicalPlan, - select_list: Vec, + select_list: Vec, ) -> Result { self.context.step(QueryBindStep::Distinct); @@ -41,82 +41,83 @@ impl> Binder<'_, '_, T, A> pub fn bind_distinct_output_exprs<'c>( &mut self, - select_list: &[ScalarExpression], - exprs: impl IntoIterator, + select_list: &[ExprRef], + exprs: impl IntoIterator, arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { - let mut binder = DistinctOutputBinder::new(select_list, arena); + let mut binder = DistinctOutputBinder::new(select_list); for expr in exprs { - binder.visit(expr)?; + binder.visit(expr, arena)?; } Ok(()) } pub fn bind_distinct_orderby_exprs( &mut self, - select_list: &[ScalarExpression], + select_list: &[ExprRef], orderby: &mut [SortField], arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { - let mut binder = DistinctOutputBinder::new(select_list, arena); + let mut binder = DistinctOutputBinder::new(select_list); for field in orderby { - field.expr = binder.output_ref(&field.expr).ok_or_else(|| { + let output = binder.output_ref(field.expr, arena).ok_or_else(|| { DatabaseError::InvalidValue(format!( "for SELECT DISTINCT, ORDER BY expressions must appear in select list: '{}'", - field.expr + field.expr.output_name(arena) )) })?; + field.expr = arena.alloc_expression(output); } Ok(()) } } -struct DistinctOutputBinder<'a, 'p> { - select_list: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, +struct DistinctOutputBinder<'a> { + select_list: &'a [ExprRef], } -impl<'a, 'p> DistinctOutputBinder<'a, 'p> { - fn new( - select_list: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, - ) -> Self { - Self { select_list, arena } +impl<'a> DistinctOutputBinder<'a> { + fn new(select_list: &'a [ExprRef]) -> Self { + Self { select_list } } - fn output_ref(&mut self, expr: &ScalarExpression) -> Option { + fn output_ref(&mut self, expr: ExprRef, arena: &mut PlanArena<'_>) -> Option { self.select_list .iter() .position(|candidate| { - candidate.eq_ignore_colref_pos(expr, self.arena) + candidate.eq_ignore_colref_pos(expr, arena) || candidate - .unpack_alias_ref() - .eq_ignore_colref_pos(expr.unpack_alias_ref(), self.arena) + .unpack_alias(arena) + .eq_ignore_colref_pos(expr.unpack_alias(arena), arena) }) .map(|position| { - let output_expr = &self.select_list[position]; - ScalarExpression::column_expr(output_expr.output_column_ref(self.arena), position) + let output_expr = self.select_list[position]; + ScalarExpression::column_expr(output_expr.output_column_ref(arena), position) }) } } -impl<'a> ExprVisitorMut<'a> for DistinctOutputBinder<'_, '_> { - fn visit(&mut self, expr: &'a mut ScalarExpression) -> Result<(), DatabaseError> { +impl ExprVisitorMut for DistinctOutputBinder<'_> { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { - expr: inner_expr, alias: crate::expression::AliasType::Name(_), - } = expr + .. + } = arena.expression(*expr) { - return self.visit(inner_expr); + return walk_mut_expr(self, expr, arena); } - if let Some(output_ref) = self.output_ref(expr) { - *expr = output_ref; + if let Some(output_ref) = self.output_ref(*expr, arena) { + *expr = arena.alloc_expression(output_ref); return Ok(()); } - walk_mut_expr(self, expr) + walk_mut_expr(self, expr, arena) } } @@ -145,37 +146,38 @@ mod tests { let left_column = test_column(&mut arena, "c1", LogicalType::Integer); let right_column = test_column(&mut arena, "c2", LogicalType::Integer); - let left_expr = ScalarExpression::column_expr(left_column, 0); - let right_expr = ScalarExpression::column_expr(right_column, 1); - let second_output = right_expr.clone(); - let select_output = ScalarExpression::Alias { - expr: Box::new(left_expr.clone()), + let left_expr = arena.alloc_expression(ScalarExpression::column_expr(left_column, 0)); + let right_expr = arena.alloc_expression(ScalarExpression::column_expr(right_column, 1)); + let second_output = right_expr; + let select_output = arena.alloc_expression(ScalarExpression::Alias { + expr: left_expr, alias: AliasType::Name("v".to_string()), - }; - let select_list = [select_output.clone(), right_expr.clone()]; + }); + let select_list = [select_output, right_expr]; - let mut order_by_alias = ScalarExpression::Alias { - expr: Box::new(left_expr), + let mut order_by_alias = arena.alloc_expression(ScalarExpression::Alias { + expr: left_expr, alias: AliasType::Name("v".to_string()), - }; + }); let mut order_by_second = right_expr; { - let mut binder = DistinctOutputBinder::new(&select_list, &mut arena); - binder.visit(&mut order_by_alias)?; - binder.visit(&mut order_by_second)?; + let mut binder = DistinctOutputBinder::new(&select_list); + binder.visit(&mut order_by_alias, &mut arena)?; + binder.visit(&mut order_by_second, &mut arena)?; } - let expected_alias = ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( - select_output.output_column_ref(&mut arena), - 0, - )), + let select_column = select_output.output_column_ref(&mut arena); + let expected_inner = + arena.alloc_expression(ScalarExpression::column_expr(select_column, 0)); + let expected_alias = arena.alloc_expression(ScalarExpression::Alias { + expr: expected_inner, alias: AliasType::Name("v".to_string()), - }; - assert!(order_by_alias.eq_ignore_colref_pos(&expected_alias, &arena)); + }); + assert!(order_by_alias.eq_ignore_colref_pos(expected_alias, &arena)); + let second_column = second_output.output_column_ref(&mut arena); let expected_second = - ScalarExpression::column_expr(second_output.output_column_ref(&mut arena), 1); - assert!(order_by_second.eq_ignore_colref_pos(&expected_second, &arena)); + arena.alloc_expression(ScalarExpression::column_expr(second_column, 1)); + assert!(order_by_second.eq_ignore_colref_pos(expected_second, &arena)); Ok(()) } @@ -185,25 +187,25 @@ mod tests { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let column = test_column(&mut arena, "c1", LogicalType::Integer); - let expr = ScalarExpression::column_expr(column, 0); - let select_output = ScalarExpression::Alias { - expr: Box::new(expr.clone()), + let expr = arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let select_output = arena.alloc_expression(ScalarExpression::Alias { + expr, alias: AliasType::Name("v".to_string()), - }; + }); - let mut target = ScalarExpression::Alias { - expr: Box::new(ScalarExpression::Constant(1_i32.into())), - alias: AliasType::Expr(Box::new(expr)), - }; + let constant = arena.alloc_expression(ScalarExpression::Constant(1_i32.into())); + let mut target = arena.alloc_expression(ScalarExpression::Alias { + expr: constant, + alias: AliasType::Expr(expr), + }); { - let mut binder = - DistinctOutputBinder::new(std::slice::from_ref(&select_output), &mut arena); - binder.visit(&mut target)?; + let mut binder = DistinctOutputBinder::new(std::slice::from_ref(&select_output)); + binder.visit(&mut target, &mut arena)?; } - let expected = - ScalarExpression::column_expr(select_output.output_column_ref(&mut arena), 0); - assert!(target.eq_ignore_colref_pos(&expected, &arena)); + let output_column = select_output.output_column_ref(&mut arena); + let expected = arena.alloc_expression(ScalarExpression::column_expr(output_column, 0)); + assert!(target.eq_ignore_colref_pos(expected, &arena)); Ok(()) } diff --git a/src/binder/expr.rs b/src/binder/expr.rs index 1febc2b4..5470c1ed 100644 --- a/src/binder/expr.rs +++ b/src/binder/expr.rs @@ -22,10 +22,10 @@ use super::{Binder, BinderContext, QueryBindStep, SubQueryType}; use crate::expression::function::scala::{ArcScalarFunctionImpl, ScalarFunction}; use crate::expression::function::table::TableFunction; use crate::expression::function::FunctionSummary; -use crate::expression::{AliasType, ScalarExpression}; +use crate::expression::{AliasType, ScalarExpression, TypeCast}; use crate::planner::operator::mark_apply::MarkApplyQuantifier; use crate::planner::operator::scalar_subquery::ScalarSubqueryOperator; -use crate::planner::{LogicalPlan, PlanArena}; +use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::value::{DataValue, Utf8Type}; use crate::types::{CharLengthUnits, LogicalType}; @@ -53,7 +53,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T fn find_column_in_scope( context: &BinderContext<'a, T>, - arena: &mut PlanArena, + arena: &PlanArena, column_name: &str, ) -> Option { let mut position_offset = 0; @@ -83,7 +83,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T ) -> Result<(ScalarExpression, LogicalPlan), DatabaseError> { let (exprs, is_tuple) = match expr { ScalarExpression::Tuple(exprs) => (exprs, true), - expr => (vec![expr], false), + expr => (vec![arena.alloc_expression(expr)], false), }; let mut alias_exprs = Vec::with_capacity(exprs.len()); let mut alias_refs = Vec::with_capacity(exprs.len()); @@ -91,10 +91,13 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T for (position, expr) in exprs.into_iter().enumerate() { let (alias_expr, alias_ref) = self.bind_temp_table_alias(expr, position, arena); if !is_tuple { - let alias_plan = Self::build_project_plan(sub_query, vec![alias_expr.clone()]); + let alias_plan = Self::build_project_plan( + sub_query, + vec![arena.alloc_expression(alias_expr.clone())], + ); return Ok((alias_expr, alias_plan)); } - alias_exprs.push(alias_expr); + alias_exprs.push(arena.alloc_expression(alias_expr)); alias_refs.push(alias_ref); } @@ -104,20 +107,21 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T pub(crate) fn bind_temp_table_alias( &mut self, - expr: ScalarExpression, + expr: ExprRef, position: usize, arena: &mut PlanArena, - ) -> (ScalarExpression, ScalarExpression) { + ) -> (ScalarExpression, ExprRef) { let output_column = expr.output_column_ref(arena); let mut alias_column = arena.clone_column(output_column); alias_column.set_ref_table(arena.temp_table(), 0, true); let alias_column = arena.alloc_column(alias_column); - let alias_ref = ScalarExpression::column_expr(alias_column, position); + let alias_ref = + arena.alloc_expression(ScalarExpression::column_expr(alias_column, position)); ( ScalarExpression::Alias { - expr: Box::new(expr), - alias: AliasType::Expr(Box::new(alias_ref.clone())), + expr, + alias: AliasType::Expr(alias_ref), }, alias_ref, ) @@ -171,7 +175,9 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T let columns = sub_query_schema .iter() .enumerate() - .map(|(position, column)| ScalarExpression::column_expr(*column, position)) + .map(|(position, column)| { + arena.alloc_expression(ScalarExpression::column_expr(*column, position)) + }) .collect::>(); ScalarExpression::Tuple(columns) } else { @@ -231,11 +237,8 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T } let (sub_query, correlated) = self.bind_subquery_plan(arena, build)?; - let (_, marker_ref) = self.bind_temp_table_alias( - ScalarExpression::Constant(DataValue::Boolean(true)), - 0, - arena, - ); + let marker = arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); + let (_, marker_ref) = self.bind_temp_table_alias(marker, 0, arena); let output_column = marker_ref.output_column_ref(arena); self.context.sub_query(SubQueryType::ExistsSubQuery { plan: sub_query, @@ -245,12 +248,12 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T if negated { Ok(ScalarExpression::Unary { op: expression::UnaryOperator::Not, - expr: Box::new(marker_ref), + expr: marker_ref, evaluator: None, ty: LogicalType::Boolean, }) } else { - Ok(marker_ref) + Ok(ScalarExpression::column_expr(output_column, 0)) } } @@ -258,7 +261,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T &mut self, quantifier: MarkApplyQuantifier, negated: bool, - left_expr: ScalarExpression, + left_expr: ExprRef, compare_op: expression::BinaryOperator, arena: &mut PlanArena<'arena>, build: F, @@ -280,18 +283,16 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T } let (alias_expr, sub_query) = self.bind_temp_table(column, sub_query, arena)?; + let alias_expr = arena.alloc_expression(alias_expr); let predicate = ScalarExpression::Binary { op: compare_op, - left_expr: Box::new(left_expr), - right_expr: Box::new(alias_expr), + left_expr, + right_expr: alias_expr, evaluator: None, ty: LogicalType::Boolean, }; - let (_, marker_ref) = self.bind_temp_table_alias( - ScalarExpression::Constant(DataValue::Boolean(true)), - 0, - arena, - ); + let marker = arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); + let (_, marker_ref) = self.bind_temp_table_alias(marker, 0, arena); let output_column = marker_ref.output_column_ref(arena); self.context.sub_query(SubQueryType::QuantifiedSubQuery { quantifier, @@ -305,12 +306,12 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T if negated { Ok(ScalarExpression::Unary { op: expression::UnaryOperator::Not, - expr: Box::new(marker_ref), + expr: marker_ref, evaluator: None, ty: LogicalType::Boolean, }) } else { - Ok(marker_ref) + Ok(ScalarExpression::column_expr(output_column, 0)) } } @@ -329,7 +330,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T .find(|((table, column), _)| table.is_none() && column == column_name) { return Ok(ScalarExpression::Alias { - expr: Box::new(expr.clone()), + expr: *expr, alias: AliasType::Name(column_name.to_string()), }); } @@ -367,9 +368,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T .get(column_name) .map(|using_column| using_column.visible_expr(arena)) .transpose()? - .or_else(|| { - Self::find_column_in_scope(context, arena, column_name) - })) + .or_else(|| Self::find_column_in_scope(context, arena, column_name))) }; let mut got_column = find_visible_column(&self.context)?; if got_column.is_none() { @@ -387,13 +386,11 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T pub(crate) fn bind_binary_op_expr( &mut self, - left_expr: ScalarExpression, - right_expr: ScalarExpression, + left_expr: ExprRef, + right_expr: ExprRef, op: expression::BinaryOperator, arena: &mut PlanArena, ) -> Result { - let left_expr = Box::new(left_expr); - let right_expr = Box::new(right_expr); let left_ty = left_expr.return_type(arena); let right_ty = right_expr.return_type(arena); let ty = match &op { @@ -439,11 +436,10 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T pub(crate) fn bind_unary_op_expr( &mut self, - expr: ScalarExpression, + expr: ExprRef, op: expression::UnaryOperator, arena: &mut PlanArena, ) -> Result { - let expr = Box::new(expr); let ty = if let expression::UnaryOperator::Not = op { LogicalType::Boolean } else { @@ -461,7 +457,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T pub(crate) fn bind_aggregate_function( &mut self, kind: AggKind, - args: Vec, + args: Vec, is_distinct: bool, arena: &mut PlanArena, ) -> Result { @@ -508,7 +504,7 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T pub(crate) fn bind_function_call( &mut self, function_name: String, - mut args: Vec, + mut args: Vec, arena: &mut PlanArena, ) -> Result { match function_name.as_str() { @@ -517,9 +513,9 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T return Err(DatabaseError::MisMatch("number of if() parameters", "3")); } let ty = Self::return_type(&args[1], &args[2], arena)?; - let right_expr = Box::new(args.pop().unwrap()); - let left_expr = Box::new(args.pop().unwrap()); - let condition = Box::new(args.pop().unwrap()); + let right_expr = args.pop().unwrap(); + let left_expr = args.pop().unwrap(); + let condition = args.pop().unwrap(); return Ok(ScalarExpression::If { condition, @@ -536,8 +532,8 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T )); } let ty = Self::return_type(&args[0], &args[1], arena)?; - let right_expr = Box::new(args.pop().unwrap()); - let left_expr = Box::new(args.pop().unwrap()); + let right_expr = args.pop().unwrap(); + let left_expr = args.pop().unwrap(); return Ok(ScalarExpression::NullIf { left_expr, @@ -553,8 +549,8 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T )); } let ty = Self::return_type(&args[0], &args[1], arena)?; - let right_expr = Box::new(args.pop().unwrap()); - let left_expr = Box::new(args.pop().unwrap()); + let right_expr = args.pop().unwrap(); + let left_expr = args.pop().unwrap(); return Ok(ScalarExpression::IfNull { left_expr, @@ -615,8 +611,8 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T } pub(crate) fn return_type( - expr_1: &ScalarExpression, - expr_2: &ScalarExpression, + expr_1: &ExprRef, + expr_2: &ExprRef, arena: &PlanArena, ) -> Result { let temp_ty_1 = expr_1.return_type(arena); diff --git a/src/binder/mod.rs b/src/binder/mod.rs index 2a87623d..edd0ca01 100644 --- a/src/binder/mod.rs +++ b/src/binder/mod.rs @@ -63,10 +63,10 @@ use crate::catalog::view::View; use crate::catalog::{ColumnRef, TableCatalog, TableName}; use crate::db::{ScalaFunctions, TableFunctions}; use crate::errors::DatabaseError; -use crate::expression::ScalarExpression; +use crate::expression::{ScalarExpression, TypeCast}; use crate::planner::operator::join::JoinType; use crate::planner::operator::mark_apply::MarkApplyQuantifier; -use crate::planner::{LogicalPlan, PlanArena, PlanRef}; +use crate::planner::{ExprRef, LogicalPlan, PlanArena, PlanRef}; use crate::storage::{TableCache, Transaction, ViewCache}; use crate::types::tuple::Schema; use crate::types::value::DataValue; @@ -206,13 +206,13 @@ impl UsingColumn { pub(crate) fn visible_expr( &self, - arena: &PlanArena, + arena: &mut PlanArena, ) -> Result { match self.join_type { JoinType::RightOuter => Ok(self.right_expr()), JoinType::Full => { - let left_expr = self.left_expr(); - let right_expr = self.right_expr(); + let left_expr = arena.alloc_expression(self.left_expr()); + let right_expr = arena.alloc_expression(self.right_expr()); let left_ty = left_expr.return_type(arena); let right_ty = right_expr.return_type(arena); let ty = LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned(); @@ -248,11 +248,11 @@ pub struct BinderContext<'a, T: Transaction> { ctes: Vec, pub(crate) cte_depth: usize, // alias - expr_aliases: BTreeMap<(Option, String), ScalarExpression>, + expr_aliases: BTreeMap<(Option, String), ExprRef>, table_aliases: HashMap, // agg - group_by_exprs: Vec, - pub(crate) agg_calls: Vec, + group_by_exprs: Vec, + pub(crate) agg_calls: Vec, // join using: HashMap, @@ -599,12 +599,7 @@ impl<'a, T: Transaction> BinderContext<'a, T> { Ok(()) } - pub fn add_alias( - &mut self, - alias_table: Option, - alias_column: String, - expr: ScalarExpression, - ) { + pub fn add_alias(&mut self, alias_table: Option, alias_column: String, expr: ExprRef) { self.expr_aliases.insert((alias_table, alias_column), expr); } @@ -612,7 +607,7 @@ impl<'a, T: Transaction> BinderContext<'a, T> { self.table_aliases.insert(alias.clone(), table.clone()); } - pub fn has_agg_call(&self, expr: &ScalarExpression) -> bool { + pub fn has_agg_call(&self, expr: &ExprRef) -> bool { self.group_by_exprs.contains(expr) } } diff --git a/src/binder/parser.rs b/src/binder/parser.rs index 42206ad0..804e025f 100644 --- a/src/binder/parser.rs +++ b/src/binder/parser.rs @@ -27,7 +27,7 @@ use crate::expression::agg::AggKind; use crate::expression::simplify::ConstantCalculator; use crate::expression::visitor_mut::ExprVisitorMut; use crate::expression::window::WindowFunctionKind; -use crate::expression::{AliasType, ScalarExpression}; +use crate::expression::{AliasType, ScalarExpression, TypeCast}; use crate::iter_ext::Itertools; use crate::parser::parse_sql; use crate::planner::operator::alter_table::change_column::{DefaultChange, NotNullChange}; @@ -37,7 +37,7 @@ use crate::planner::operator::project::ProjectOperator; use crate::planner::operator::recursive_cte::{RecursiveCteOperator, RecursiveScanOperator}; use crate::planner::operator::sort::SortField; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan, PlanArena}; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::storage::{Storage, Transaction}; use crate::types::value::{DataValue, Utf8Type}; use crate::types::{CharLengthUnits, ColumnId, LogicalType}; @@ -336,22 +336,22 @@ struct BindStatementComplete { plan: LogicalPlan, } -struct UpdateExprTargetRemapper<'a, 'p> { +struct UpdateExprTargetRemapper<'a> { target_schema: &'a [ColumnRef], - arena: &'a PlanArena<'p>, } -impl ExprVisitorMut<'_> for UpdateExprTargetRemapper<'_, '_> { +impl ExprVisitorMut for UpdateExprTargetRemapper<'_> { fn visit_column_ref( &mut self, column: &mut ColumnRef, position: &mut usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { let Some(target_position) = self .target_schema .iter() .copied() - .position(|target_column| self.arena.same_column(target_column, *column)) + .position(|target_column| arena.same_column(target_column, *column)) else { return Err(DatabaseError::UnsupportedStmt( "joined UPDATE SET expressions can only reference target table columns".to_string(), @@ -669,15 +669,18 @@ where &mut self, expr: Expr, ty: &LogicalType, - ) -> Result { - let mut expr = self.binder.bind_expr(&expr, self.arena)?; + ) -> Result { + let mut expr = self + .binder + .bind_expr(&expr, self.arena) + .map(|expr| self.arena.alloc_expression(expr))?; if expr.any_referenced_column(self.arena, |_, _| true)? { return Err(DatabaseError::UnsupportedStmt( "column is not allowed to exist in default".to_string(), )); } - expr = ScalarExpression::type_cast(expr, Cow::Borrowed(ty), self.arena)?; + expr = expr.type_cast(Cow::Borrowed(ty), self.arena)?; Ok(expr) } @@ -784,18 +787,18 @@ where } ColumnOption::Unique(_) => column_desc.set_unique(), ColumnOption::Default(expr) => { - let mut expr = self.binder.bind_expr(&expr, self.arena)?; + let mut expr = self + .binder + .bind_expr(&expr, self.arena) + .map(|expr| self.arena.alloc_expression(expr))?; if expr.any_referenced_column(self.arena, |_, _| true)? { return Err(DatabaseError::UnsupportedStmt( "column is not allowed to exist in `default`".to_string(), )); } - expr = ScalarExpression::type_cast( - expr, - Cow::Borrowed(&column_desc.column_datatype), - self.arena, - )?; + expr = + expr.type_cast(Cow::Borrowed(&column_desc.column_datatype), self.arena)?; column_desc.default = Some(expr); } option => { @@ -848,14 +851,14 @@ where let mut columns = Vec::with_capacity(create.columns.len()); for index_column in create.columns { - match self + let expr = self .binder - .bind_expr(&index_column.column.expr, self.arena)? - { - ScalarExpression::ColumnRef { column, .. } => columns.push(column), - expr => { + .bind_expr(&index_column.column.expr, self.arena)?; + match &expr { + ScalarExpression::ColumnRef { column, .. } => columns.push(*column), + _ => { return Err(DatabaseError::UnsupportedStmt(format!( - "'CREATE INDEX' by {expr}" + "'CREATE INDEX' by {expr:?}" ))) } } @@ -1008,11 +1011,12 @@ where } else { let mut columns = Vec::with_capacity(idents.len()); for ident in idents { - match self.binder.bind_column_ref_from_identifiers( + let expr = self.binder.bind_column_ref_from_identifiers( slice::from_ref(ident), Some(table_name.as_ref()), self.arena, - )? { + )?; + match expr { ScalarExpression::ColumnRef { column, .. } => columns.push(column), _ => return Err(DatabaseError::UnsupportedStmt(ident.to_string())), } @@ -1032,9 +1036,14 @@ where for (i, expr) in expr_row.iter().enumerate() { let span = expr.span(); - let mut expression = self.binder.bind_expr(expr, self.arena)?; + let mut expr_ref = self + .binder + .bind_expr(expr, self.arena) + .map(|expr| self.arena.alloc_expression(expr))?; - ConstantCalculator::new(self.arena).visit(&mut expression)?; + ConstantCalculator::new(self.arena).visit(&mut expr_ref, self.arena)?; + let expression = + std::mem::replace(self.arena.expression_mut(expr_ref), ScalarExpression::Empty); match expression { ScalarExpression::Constant(mut value) => { let column = self.arena.column(schema_ref[i]); @@ -1054,7 +1063,7 @@ where ScalarExpression::Empty => { let column = self.arena.column(schema_ref[i]); let default_value = column - .default_value()? + .default_value(self.arena)? .ok_or(DatabaseError::DefaultNotExist)?; if default_value.is_null() && !column.nullable() { return Err(attach_span_if_absent( @@ -1117,12 +1126,17 @@ where .iter() .copied() .enumerate() - .map(|(position, target_column)| ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( + .map(|(position, target_column)| { + let expr = self.arena.alloc_expression(ScalarExpression::column_expr( input_schema[position], position, - )), - alias: AliasType::Name(self.arena.column(target_column).name().to_string()), + )); + self.arena.alloc_expression(ScalarExpression::Alias { + expr, + alias: AliasType::Name( + self.arena.column(target_column).name().to_string(), + ), + }) }) .collect::>() } else { @@ -1138,13 +1152,14 @@ where ident.span, ) })?; - projection.push(ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( - input_schema[position], - position, - )), + let expr = self.arena.alloc_expression(ScalarExpression::column_expr( + input_schema[position], + position, + )); + projection.push(self.arena.alloc_expression(ScalarExpression::Alias { + expr, alias: AliasType::Name(self.arena.column(column).name().to_string()), - }); + })); } projection } @@ -1183,17 +1198,20 @@ where } for Assignment { target, value } in &update.assignments { let expression = self.binder.bind_expr(value, self.arena)?; - let mut bind_assignment = |name: &ObjectName, + let mut bind_assignment = |binder: &mut Binder<'_, '_, T, A>, + arena: &mut PlanArena, + name: &ObjectName, expression: ScalarExpression| -> Result<(), DatabaseError> { let ident = single_ident_from_object_name(name)?; let column = { - match self.binder.bind_column_ref_from_identifiers( + let expr = binder.bind_column_ref_from_identifiers( slice::from_ref(&ident), Some(table_name.as_ref()), - self.arena, - )? { + arena, + )?; + match expr { ScalarExpression::ColumnRef { column, .. } => column, _ => { return Err(attach_span_if_absent( @@ -1204,34 +1222,32 @@ where } }; - let mut expr = if matches!(expression, ScalarExpression::Empty) { - let column_catalog = self.arena.column(column); - let default_value = column_catalog - .default_value()? - .ok_or(DatabaseError::DefaultNotExist)?; - ScalarExpression::Constant(default_value) - } else { - expression + let mut expr = match expression { + ScalarExpression::Empty => { + let column_catalog = arena.column(column); + let default_value = column_catalog + .default_value(arena)? + .ok_or(DatabaseError::DefaultNotExist)?; + arena.alloc_expression(ScalarExpression::Constant(default_value)) + } + expression => arena.alloc_expression(expression), }; - let column_catalog = self.arena.column(column); - expr = ScalarExpression::type_cast( - expr, - Cow::Borrowed(column_catalog.datatype()), - self.arena, - )?; + let datatype = arena.column(column).datatype().clone(); + expr = expr.type_cast(Cow::Owned(datatype), arena)?; if is_joined_update { UpdateExprTargetRemapper { target_schema: &target_schema, - arena: self.arena, } - .visit(&mut expr)?; + .visit(&mut expr, arena)?; } value_exprs.push((column, expr)); Ok(()) }; match target { - AssignmentTarget::ColumnName(name) => bind_assignment(name, expression)?, + AssignmentTarget::ColumnName(name) => { + bind_assignment(self.binder, self.arena, name, expression)? + } AssignmentTarget::Tuple(names) => { let expected = names.len(); let ScalarExpression::Tuple(exprs) = expression else { @@ -1244,7 +1260,11 @@ where loop { match (names.next(), exprs.next()) { (Some(name), Some(expression)) => { - bind_assignment(name, expression)? + let expression = std::mem::replace( + self.arena.expression_mut(expression), + ScalarExpression::Empty, + ); + bind_assignment(self.binder, self.arena, name, expression)? } (None, None) => break, _ => return Err(DatabaseError::ValuesLenMismatch(expected, got)), @@ -1260,7 +1280,10 @@ where .copied() .enumerate() .map(|(index, column)| { - ScalarExpression::column_expr(column, target_offset + index) + self.arena.alloc_expression(ScalarExpression::column_expr( + column, + target_offset + index, + )) }) .collect(); plan = LogicalPlan::new( @@ -1406,7 +1429,8 @@ where ) -> Result, DatabaseError> { let predicate = if let Some(predicate) = selection { Some(with_query_bind_step!(self.binder, QueryBindStep::Where, { - self.binder.bind_expr(predicate, self.arena)? + let predicate = self.binder.bind_expr(predicate, self.arena)?; + self.arena.alloc_expression(predicate) })?) } else { None @@ -1437,7 +1461,11 @@ where } group_by_exprs .iter() - .map(|expr| self.binder.bind_expr(expr, self.arena)) + .map(|expr| { + self.binder + .bind_expr(expr, self.arena) + .map(|expr| self.arena.alloc_expression(expr)) + }) .collect::, DatabaseError>>()? } GroupByExpr::All(_) => { @@ -1450,15 +1478,17 @@ where let having = having .map(|having| { with_query_bind_step!(self.binder, QueryBindStep::Having, { - self.binder.bind_expr(having, self.arena)? + let having = self.binder.bind_expr(having, self.arena)?; + self.arena.alloc_expression(having) }) }) .transpose()?; self.aggregate(group_by, having, orderby, |binder, arena, orderby| { let OrderByExpr { expr, options, .. } = orderby; with_query_bind_step!(binder, QueryBindStep::Sort, { + let expr = binder.bind_expr(expr, arena)?; SortField::new( - binder.bind_expr(expr, arena)?, + arena.alloc_expression(expr), options.asc.is_none_or(|asc| asc), options.nulls_first.unwrap_or(false), ) @@ -2028,7 +2058,10 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< arena: &mut PlanArena, ) -> Result { match constraint { - JoinConstraint::On(expr) => Ok(JoinConstraintInput::On(self.bind_expr(expr, arena)?)), + JoinConstraint::On(expr) => { + let expr = self.bind_expr(expr, arena)?; + Ok(JoinConstraintInput::On(arena.alloc_expression(expr))) + } JoinConstraint::Using(names) => Ok(JoinConstraintInput::Using( names .iter() @@ -2077,8 +2110,12 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< self.bind_column_ref_from_identifiers(idents, None, arena) } Expr::BinaryOp { left, right, op } => { - let left_expr = self.bind_expr(left, arena)?; - let right_expr = self.bind_expr(right, arena)?; + let left_expr = self + .bind_expr(left, arena) + .map(|expr| arena.alloc_expression(expr))?; + let right_expr = self + .bind_expr(right, arena) + .map(|expr| arena.alloc_expression(expr))?; self.bind_binary_op_expr(left_expr, right_expr, op.clone().try_into()?, arena) } Expr::Value(v) => { @@ -2103,7 +2140,9 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< Expr::Function(func) => self.bind_function_sql(func, arena), Expr::Nested(expr) => self.bind_expr(expr, arena), Expr::UnaryOp { expr, op } => { - let expr = self.bind_expr(expr, arena)?; + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; self.bind_unary_op_expr(expr, (*op).try_into()?, arena) } Expr::Like { @@ -2113,8 +2152,12 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< escape_char, any: _, } => { - let left_expr = Box::new(self.bind_expr(expr, arena)?); - let right_expr = Box::new(self.bind_expr(pattern, arena)?); + let left_expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + let right_expr = self + .bind_expr(pattern, arena) + .map(|expr| arena.alloc_expression(expr))?; let escape_char = Self::parse_like_escape_char(escape_char)?; let op = if *negated { expression::BinaryOperator::NotLike(escape_char) @@ -2129,14 +2172,24 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< ty: LogicalType::Boolean, }) } - Expr::IsNull(expr) => Ok(ScalarExpression::IsNull { - negated: false, - expr: Box::new(self.bind_expr(expr, arena)?), - }), - Expr::IsNotNull(expr) => Ok(ScalarExpression::IsNull { - negated: true, - expr: Box::new(self.bind_expr(expr, arena)?), - }), + Expr::IsNull(expr) => { + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + Ok(ScalarExpression::IsNull { + negated: false, + expr, + }) + } + Expr::IsNotNull(expr) => { + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + Ok(ScalarExpression::IsNull { + negated: true, + expr, + }) + } Expr::InList { expr, list, @@ -2144,21 +2197,25 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< } => { let args = list .iter() - .map(|expr| self.bind_expr(expr, arena)) + .map(|expr| { + self.bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr)) + }) .try_collect()?; + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; Ok(ScalarExpression::In { negated: *negated, - expr: Box::new(self.bind_expr(expr, arena)?), + expr, args, }) } Expr::Cast { expr, data_type, .. - } => ScalarExpression::type_cast( - self.bind_expr(expr, arena)?, - Cow::Owned(LogicalType::try_from(data_type.clone())?), - arena, - ), + } => self + .bind_expr(expr, arena)? + .type_cast(Cow::Owned(LogicalType::try_from(data_type.clone())?), arena), Expr::TypedString(TypedString { data_type, value, .. }) => { @@ -2181,12 +2238,23 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< negated, low, high, - } => Ok(ScalarExpression::Between { - negated: *negated, - expr: Box::new(self.bind_expr(expr, arena)?), - left_expr: Box::new(self.bind_expr(low, arena)?), - right_expr: Box::new(self.bind_expr(high, arena)?), - }), + } => { + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + let left_expr = self + .bind_expr(low, arena) + .map(|expr| arena.alloc_expression(expr))?; + let right_expr = self + .bind_expr(high, arena) + .map(|expr| arena.alloc_expression(expr))?; + Ok(ScalarExpression::Between { + negated: *negated, + expr, + left_expr, + right_expr, + }) + } Expr::Substring { expr, substring_for, @@ -2197,22 +2265,32 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let mut from_expr = None; if let Some(expr) = substring_for { - for_expr = Some(Box::new(self.bind_expr(expr, arena)?)) + let expr = self.bind_expr(expr, arena)?; + for_expr = Some(arena.alloc_expression(expr)) } if let Some(expr) = substring_from { - from_expr = Some(Box::new(self.bind_expr(expr, arena)?)) + let expr = self.bind_expr(expr, arena)?; + from_expr = Some(arena.alloc_expression(expr)) } + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; Ok(ScalarExpression::SubString { - expr: Box::new(self.bind_expr(expr, arena)?), + expr, for_expr, from_expr, }) } - Expr::Position { expr, r#in } => Ok(ScalarExpression::Position { - expr: Box::new(self.bind_expr(expr, arena)?), - in_expr: Box::new(self.bind_expr(r#in, arena)?), - }), + Expr::Position { expr, r#in } => { + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + let in_expr = self + .bind_expr(r#in, arena) + .map(|expr| arena.alloc_expression(expr))?; + Ok(ScalarExpression::Position { expr, in_expr }) + } Expr::Trim { expr, trim_what, @@ -2221,10 +2299,14 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< } => { let mut trim_what_expr = None; if let Some(trim_what) = trim_what { - trim_what_expr = Some(Box::new(self.bind_expr(trim_what, arena)?)) + let trim_what = self.bind_expr(trim_what, arena)?; + trim_what_expr = Some(arena.alloc_expression(trim_what)) } + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; Ok(ScalarExpression::Trim { - expr: Box::new(self.bind_expr(expr, arena)?), + expr, trim_what_expr, trim_where: (*trim_where).map(Into::into), }) @@ -2253,7 +2335,8 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let mut bound_exprs = Vec::with_capacity(exprs.len()); for expr in exprs { - bound_exprs.push(self.bind_expr(expr, arena)?); + let expr = self.bind_expr(expr, arena)?; + bound_exprs.push(arena.alloc_expression(expr)); } Ok(ScalarExpression::Tuple(bound_exprs)) } @@ -2277,20 +2360,28 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let mut operand_expr = None; let mut ty = LogicalType::SqlNull; if let Some(expr) = operand { - operand_expr = Some(Box::new(self.bind_expr(expr, arena)?)); + let expr = self.bind_expr(expr, arena)?; + operand_expr = Some(arena.alloc_expression(expr)); } let mut expr_pairs = Vec::with_capacity(conditions.len()); for when in conditions { - let result = self.bind_expr(&when.result, arena)?; + let result = self + .bind_expr(&when.result, arena) + .map(|expr| arena.alloc_expression(expr))?; let result_ty = result.return_type(arena).into_owned(); fn_check_ty(&mut ty, result_ty)?; - expr_pairs.push((self.bind_expr(&when.condition, arena)?, result)) + let condition = self + .bind_expr(&when.condition, arena) + .map(|expr| arena.alloc_expression(expr))?; + expr_pairs.push((condition, result)) } let mut else_expr = None; if let Some(expr) = else_result { - let temp_expr = Box::new(self.bind_expr(expr, arena)?); + let temp_expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; let else_ty = temp_expr.return_type(arena).into_owned(); fn_check_ty(&mut ty, else_ty)?; @@ -2345,7 +2436,9 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< subquery: &Query, arena: &mut PlanArena, ) -> Result { - let left_expr = self.bind_expr(expr, arena)?; + let left_expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; self.bind_quantified_subquery_plan( quantifier, negated, @@ -2426,8 +2519,13 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< FunctionArg::Unnamed(arg) => arg, }; match arg_expr { - FunctionArgExpr::Expr(expr) => args.push(self.bind_expr(expr, arena)?), - FunctionArgExpr::Wildcard => args.push(Self::wildcard_expr()), + FunctionArgExpr::Expr(expr) => { + let expr = self.bind_expr(expr, arena)?; + args.push(arena.alloc_expression(expr)) + } + FunctionArgExpr::Wildcard => { + args.push(arena.alloc_expression(Self::wildcard_expr())) + } expr => { return Err(DatabaseError::UnsupportedStmt(format!( "function arg: {expr:#?}" @@ -2468,7 +2566,7 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< fn bind_window_call( &mut self, kind: WindowFunctionKind, - args: Vec, + args: Vec, is_distinct: bool, over: &WindowType, arena: &mut PlanArena, @@ -2505,14 +2603,18 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let partition_by = spec .partition_by .iter() - .map(|expr| self.bind_expr(expr, arena)) + .map(|expr| { + self.bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr)) + }) .collect::, _>>()?; let order_by = spec .order_by .iter() .map(|OrderByExpr { expr, options, .. }| { + let expr = self.bind_expr(expr, arena)?; Ok(SortField::new( - self.bind_expr(expr, arena)?, + arena.alloc_expression(expr), options.asc.unwrap_or(true), options.nulls_first.unwrap_or(false), )) @@ -2673,9 +2775,13 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let mut row = Vec::with_capacity(values_len); for (col_index, expr) in expr_row.iter().enumerate() { - let mut expression = self.bind_expr(expr, arena)?; - ConstantCalculator::new(arena).visit(&mut expression)?; + let mut expression = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; + ConstantCalculator::new(arena).visit(&mut expression, arena)?; + let expression = + std::mem::replace(arena.expression_mut(expression), ScalarExpression::Empty); if let ScalarExpression::Constant(value) = expression { let value_type = value.logical_type(); @@ -2726,19 +2832,25 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< self.context.add_alias( None, arena.column(*column).name().to_string(), - ScalarExpression::column_expr(*column, position), + arena.alloc_expression(ScalarExpression::column_expr(*column, position)), ); } let sort_fields = self - .extract_having_orderby_aggregate_exprs(None, Some(orderbys), |binder, orderby| { - let OrderByExpr { expr, options, .. } = orderby; - Ok(SortField::new( - binder.bind_expr(expr, arena)?, - options.asc.is_none_or(|asc| asc), - options.nulls_first.unwrap_or(false), - )) - })? + .extract_having_orderby_aggregate_exprs( + None, + Some(orderbys), + |binder, orderby, arena| { + let OrderByExpr { expr, options, .. } = orderby; + let expr = binder.bind_expr(expr, arena)?; + Ok(SortField::new( + arena.alloc_expression(expr), + options.asc.is_none_or(|asc| asc), + options.nulls_first.unwrap_or(false), + )) + }, + arena, + )? .1; self.context.expr_aliases = saved_aliases; @@ -2755,7 +2867,8 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< arena: &mut PlanArena, ) -> Result { let predicate = with_query_bind_step!(self, QueryBindStep::Where, { - self.bind_expr(predicate, arena)? + let predicate = self.bind_expr(predicate, arena)?; + arena.alloc_expression(predicate) })?; self.bind_where_expr(children, predicate, arena) @@ -2765,23 +2878,27 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< &mut self, items: &[SelectItem], arena: &mut PlanArena, - ) -> Result, DatabaseError> { + ) -> Result, DatabaseError> { let mut select_items = vec![]; for item in items { match item { - SelectItem::UnnamedExpr(expr) => select_items.push(self.bind_expr(expr, arena)?), - SelectItem::ExprWithAlias { expr, alias } => { + SelectItem::UnnamedExpr(expr) => { let expr = self.bind_expr(expr, arena)?; + select_items.push(arena.alloc_expression(expr)); + } + SelectItem::ExprWithAlias { expr, alias } => { + let expr = self + .bind_expr(expr, arena) + .map(|expr| arena.alloc_expression(expr))?; let alias_name = lower_ident(alias).into_owned(); - self.context - .add_alias(None, alias_name.clone(), expr.clone()); + self.context.add_alias(None, alias_name.clone(), expr); - select_items.push(ScalarExpression::Alias { - expr: Box::new(expr), + select_items.push(arena.alloc_expression(ScalarExpression::Alias { + expr, alias: AliasType::Name(alias_name), - }); + })); } SelectItem::Wildcard(_) => { let visible_names = self @@ -3052,8 +3169,8 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< ) -> Result { let span = expr.span(); let bound_expr = self.bind_expr(expr, arena)?; - match bound_expr { - ScalarExpression::Constant(dv) => match &dv { + match &bound_expr { + ScalarExpression::Constant(dv) => match dv { DataValue::Int32(v) if *v >= 0 => Ok(*v as usize), DataValue::Int64(v) if *v >= 0 => Ok(*v as usize), _ => Err(DatabaseError::InvalidType), diff --git a/src/binder/select.rs b/src/binder/select.rs index 1bad19eb..ac358302 100644 --- a/src/binder/select.rs +++ b/src/binder/select.rs @@ -24,7 +24,7 @@ use crate::{ }, types::value::DataValue, }; -use std::collections::HashSet; +use std::{borrow::Cow, collections::HashSet}; use super::{Binder, BinderContext, QueryBindStep, SetOperatorKind, Source, SubQueryType}; @@ -32,7 +32,7 @@ use crate::catalog::{ColumnRef, ColumnRelation, TableName}; use crate::errors::DatabaseError; use crate::execution::dql::join::joins_nullable; use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut, PositionShift}; -use crate::expression::{AliasType, BinaryOperator}; +use crate::expression::{AliasType, BinaryOperator, TypeCast}; use crate::iter_ext::Itertools; use crate::planner::operator::function_scan::FunctionScanOperator; use crate::planner::operator::insert::InsertOperator; @@ -40,27 +40,27 @@ use crate::planner::operator::join::JoinCondition; use crate::planner::operator::set_membership::{SetMembershipKind, SetMembershipOperator}; use crate::planner::operator::sort::{SortField, SortOperator}; use crate::planner::operator::union::UnionOperator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::tuple::Schema; use crate::types::{ColumnId, LogicalType}; -struct RightSidePositionGlobalizer<'a, 'p> { +struct RightSidePositionGlobalizer<'a> { right_schema: &'a Schema, left_len: usize, - arena: &'a crate::planner::PlanArena<'p>, } -impl<'a> ExprVisitorMut<'a> for RightSidePositionGlobalizer<'_, '_> { +impl ExprVisitorMut for RightSidePositionGlobalizer<'_> { fn visit_column_ref( &mut self, - column: &'a mut ColumnRef, - position: &'a mut usize, + column: &mut ColumnRef, + position: &mut usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if self .right_schema .iter() - .any(|right| self.arena.same_column(*right, *column)) + .any(|right| arena.same_column(*right, *column)) { *position += self.left_len; } @@ -74,28 +74,28 @@ struct AppendedRightOutput { output_position: usize, } -struct SplitScopePositionRebinder<'a, 'p> { +struct SplitScopePositionRebinder<'a> { left_schema: &'a Schema, right_schema: &'a Schema, - arena: &'a crate::planner::PlanArena<'p>, } -impl ExprVisitorMut<'_> for SplitScopePositionRebinder<'_, '_> { +impl ExprVisitorMut for SplitScopePositionRebinder<'_> { fn visit_column_ref( &mut self, column: &mut ColumnRef, position: &mut usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if let Some(left_position) = self .left_schema .iter() - .position(|candidate| self.arena.same_column(*candidate, *column)) + .position(|candidate| arena.same_column(*candidate, *column)) { *position = left_position; } else if let Some(right_position) = self .right_schema .iter() - .position(|candidate| self.arena.same_column(*candidate, *column)) + .position(|candidate| arena.same_column(*candidate, *column)) { *position = right_position; } @@ -103,64 +103,61 @@ impl ExprVisitorMut<'_> for SplitScopePositionRebinder<'_, '_> { } } -struct MarkerPositionGlobalizer<'a, 'p> { +struct MarkerPositionGlobalizer<'a> { output_column: &'a ColumnRef, left_len: usize, - arena: &'a crate::planner::PlanArena<'p>, } -impl ExprVisitorMut<'_> for MarkerPositionGlobalizer<'_, '_> { +impl ExprVisitorMut for MarkerPositionGlobalizer<'_> { fn visit_column_ref( &mut self, column: &mut ColumnRef, position: &mut usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - if self.arena.same_column(*column, *self.output_column) { + if arena.same_column(*column, *self.output_column) { *position = self.left_len; } Ok(()) } } -struct ProjectionOutputBinder<'a, 'p> { - project_exprs: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, +struct ProjectionOutputBinder<'a> { + project_exprs: &'a [ExprRef], } -impl<'a, 'p> ProjectionOutputBinder<'a, 'p> { - fn new( - project_exprs: &'a [ScalarExpression], - arena: &'a mut crate::planner::PlanArena<'p>, - ) -> Self { - Self { - project_exprs, - arena, - } +impl<'a> ProjectionOutputBinder<'a> { + fn new(project_exprs: &'a [ExprRef]) -> Self { + Self { project_exprs } } - fn output_ref(&mut self, expr: &ScalarExpression) -> Option { + fn output_ref(&mut self, expr: ExprRef, arena: &mut PlanArena<'_>) -> Option { self.project_exprs .iter() .position(|candidate| { - candidate.eq_ignore_colref_pos(expr, self.arena) + candidate.eq_ignore_colref_pos(expr, arena) || candidate - .unpack_alias_ref() - .eq_ignore_colref_pos(expr.unpack_alias_ref(), self.arena) + .unpack_alias(arena) + .eq_ignore_colref_pos(expr.unpack_alias(arena), arena) }) .map(|position| { - let output_expr = &self.project_exprs[position]; - ScalarExpression::column_expr(output_expr.output_column_ref(self.arena), position) + let output_expr = self.project_exprs[position]; + ScalarExpression::column_expr(output_expr.output_column_ref(arena), position) }) } } -impl<'a> ExprVisitorMut<'a> for ProjectionOutputBinder<'_, '_> { - fn visit(&mut self, expr: &'a mut ScalarExpression) -> Result<(), DatabaseError> { - if let Some(output_ref) = self.output_ref(expr) { - *expr = output_ref; +impl ExprVisitorMut for ProjectionOutputBinder<'_> { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + if let Some(output_ref) = self.output_ref(*expr, arena) { + *expr = arena.alloc_expression(output_ref); return Ok(()); } - walk_mut_expr(self, expr) + walk_mut_expr(self, expr, arena) } } @@ -192,7 +189,7 @@ where pub(crate) binder: &'s mut Binder<'a, 'b, T, A>, pub(crate) arena: &'s mut crate::planner::PlanArena<'arena>, pub(super) plan: LogicalPlan, - pub(super) select_list: Vec, + pub(super) select_list: Vec, pub(crate) _marker: std::marker::PhantomData, } @@ -204,7 +201,7 @@ where pub(super) binder: &'s mut Binder<'a, 'b, T, A>, pub(super) arena: &'s mut crate::planner::PlanArena<'arena>, pub(super) plan: LogicalPlan, - pub(super) select_list: Vec, + pub(super) select_list: Vec, } pub(crate) struct BindPlanAggregated<'s, 'a, 'b, 'arena, T, A> @@ -215,8 +212,8 @@ where binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, plan: LogicalPlan, - select_list: Vec, - having: Option, + select_list: Vec, + having: Option, orderby: Option>, } @@ -228,7 +225,7 @@ where binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, plan: LogicalPlan, - select_list: Vec, + select_list: Vec, orderby: Option>, } @@ -240,7 +237,7 @@ where binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, plan: LogicalPlan, - select_list: Vec, + select_list: Vec, orderby: Option>, } @@ -252,7 +249,7 @@ where binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, plan: LogicalPlan, - select_list: Vec, + select_list: Vec, orderby: Option>, } @@ -264,7 +261,7 @@ where binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, plan: LogicalPlan, - select_list: Vec, + select_list: Vec, } pub(crate) struct BindPlanProjected<'s, 'a, 'b, 'arena, T, A> @@ -286,7 +283,7 @@ pub(crate) struct TableAliasInput { } pub(crate) enum JoinConstraintInput { - On(ScalarExpression), + On(ExprRef), Using(Vec), Natural, None, @@ -308,10 +305,7 @@ where } #[cfg(feature = "orm")] - pub(crate) fn filter_expr( - mut self, - predicate: ScalarExpression, - ) -> Result { + pub(crate) fn filter_expr(mut self, predicate: ExprRef) -> Result { self.plan = self .binder .bind_where_expr(self.plan, predicate, self.arena)?; @@ -335,7 +329,7 @@ where pub(crate) fn select_list( self, - select_list: Vec, + select_list: Vec, ) -> BindPlanSelectList<'s, 'a, 'b, 'arena, T, A, M> { BindPlanSelectList { binder: self.binder, @@ -353,13 +347,13 @@ where A: AsRef<[(&'static str, DataValue)]>, { #[cfg(feature = "orm")] - pub(crate) fn set_select_list(mut self, select_list: Vec) -> Self { + pub(crate) fn set_select_list(mut self, select_list: Vec) -> Self { self.select_list = select_list; self } #[cfg(feature = "orm")] - pub(crate) fn group_by_expr(self, expr: ScalarExpression) -> Result { + pub(crate) fn group_by_expr(self, expr: ExprRef) -> Result { let sorted = self .filter_expr(None)? .aggregate( @@ -405,7 +399,7 @@ where } #[cfg(feature = "orm")] - pub(crate) fn having_expr(mut self, expr: ScalarExpression) -> Result { + pub(crate) fn having_expr(mut self, expr: ExprRef) -> Result { self.plan = self.binder.bind_having(self.plan, expr, self.arena)?; Ok(self) } @@ -447,7 +441,7 @@ where #[cfg(feature = "orm")] pub fn finish(self) -> Result { for expr in &self.select_list { - if expr.has_agg_call()? || expr.has_window_call()? { + if expr.has_agg_call(self.arena)? || expr.has_window_call(self.arena)? { return self.aggregate_without_group()?.finish(); } } @@ -482,7 +476,7 @@ where { pub(crate) fn filter_expr( mut self, - predicate: Option, + predicate: Option, ) -> Result, DatabaseError> { if let Some(predicate) = predicate { self.plan = self @@ -506,8 +500,8 @@ where { pub(crate) fn aggregate( mut self, - group_by: Vec, - having: Option, + group_by: Vec, + having: Option, orderby: Option>, mut bind_sort_field: impl FnMut( &mut Binder<'a, 'b, T, A>, @@ -518,11 +512,14 @@ where self.binder .extract_select_join(&mut self.select_list, self.arena); self.binder - .extract_select_aggregate(&mut self.select_list)?; + .extract_select_aggregate(&mut self.select_list, self.arena)?; if !group_by.is_empty() { - self.binder - .extract_group_by_aggregate_exprs(&mut self.select_list, group_by)?; + self.binder.extract_group_by_aggregate_exprs( + &mut self.select_list, + group_by, + self.arena, + )?; } let mut having_orderby = (None, None); @@ -530,7 +527,8 @@ where having_orderby = self.binder.extract_having_orderby_aggregate_exprs( having, orderby, - |binder, orderby| bind_sort_field(binder, self.arena, orderby), + |binder, orderby, arena| bind_sort_field(binder, arena, orderby), + self.arena, )?; } if !self.binder.context.agg_calls.is_empty() @@ -733,19 +731,16 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } } - fn is_temp_alias_projection( - exprs: &[ScalarExpression], - arena: &crate::planner::PlanArena, - ) -> bool { + fn is_temp_alias_projection(exprs: &[ExprRef], arena: &crate::planner::PlanArena) -> bool { !exprs.is_empty() && exprs.iter().all(|expr| { matches!( - expr, + arena.expression(*expr), ScalarExpression::Alias { alias: AliasType::Expr(alias_expr), .. } if matches!( - alias_expr.unpack_alias_ref(), + alias_expr.unpack_alias_ref(arena), ScalarExpression::ColumnRef { column, .. } if matches!( &arena.column(*column).summary().relation, @@ -795,6 +790,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn localize_join_condition_from_join_scope( join_condition: &mut JoinCondition, left_len: usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { let JoinCondition::On { on, .. } = join_condition else { return Ok(()); @@ -804,7 +800,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' delta: -(left_len as isize), }; for (_, right_expr) in on { - right_shift.visit(right_expr)?; + right_shift.visit(right_expr, arena)?; } Ok(()) @@ -814,7 +810,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' join_condition: &mut JoinCondition, left_len: usize, right_schema: &Schema, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { let JoinCondition::On { filter, .. } = join_condition else { return Ok(()); @@ -824,33 +820,31 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' RightSidePositionGlobalizer { right_schema, left_len, - arena, } - .visit(expr)?; + .visit(expr, arena)?; } Ok(()) } fn localize_appended_right_outputs<'expr>( - exprs: impl Iterator, + exprs: impl Iterator, appended_outputs: &[AppendedRightOutput], - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { - struct AppendedRightOutputBinder<'a, 'p> { + struct AppendedRightOutputBinder<'a> { appended_outputs: &'a [AppendedRightOutput], - arena: &'a crate::planner::PlanArena<'p>, } - impl ExprVisitorMut<'_> for AppendedRightOutputBinder<'_, '_> { + impl ExprVisitorMut for AppendedRightOutputBinder<'_> { fn visit_column_ref( &mut self, column: &mut ColumnRef, position: &mut usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if let Some(output) = self.appended_outputs.iter().find(|output| { - *position == output.child_position - && self.arena.same_column(*column, output.column) + *position == output.child_position && arena.same_column(*column, output.column) }) { *position = output.output_position; } @@ -858,29 +852,25 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } } - let mut binder = AppendedRightOutputBinder { - appended_outputs, - arena, - }; + let mut binder = AppendedRightOutputBinder { appended_outputs }; for expr in exprs { - binder.visit(expr)?; + binder.visit(expr, arena)?; } Ok(()) } fn rebind_split_scope_positions( - expr: &mut ScalarExpression, + mut expr: ExprRef, left_schema: &Schema, right_schema: &Schema, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { SplitScopePositionRebinder { left_schema, right_schema, - arena, } - .visit(expr) + .visit(&mut expr, arena) } fn build_join_from_split_scope_predicates( @@ -888,7 +878,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' mut children: LogicalPlan, mut plan: LogicalPlan, join_ty: JoinType, - predicates: impl IntoIterator, + predicates: impl IntoIterator, rebind_positions: bool, arena: &mut crate::planner::PlanArena, ) -> Result { @@ -897,28 +887,23 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' let mut on_keys = Vec::new(); let mut filter = Vec::new(); - for mut predicate in predicates { + for predicate in predicates { if rebind_positions { - Self::rebind_split_scope_positions( - &mut predicate, - left_schema, - right_schema, - arena, - )?; + Self::rebind_split_scope_positions(predicate, &left_schema, &right_schema, arena)?; } Self::extract_join_keys( predicate, &mut on_keys, &mut filter, - left_schema, - right_schema, + &left_schema, + &right_schema, arena, )?; } let mut join_condition = JoinCondition::On { on: on_keys, - filter: Self::combine_conjuncts(filter), + filter: Self::combine_conjuncts(filter, arena), }; Self::globalize_join_filter_from_split_scope( &mut join_condition, @@ -951,28 +936,19 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' for (position, (left_schema, right_schema)) in left_schema.iter().zip(right_schema.iter()).enumerate() { - let left_column = arena.column(*left_schema); - let right_column = arena.column(*right_schema); - let cast_type = - LogicalType::max_logical_type(left_column.datatype(), right_column.datatype())?; - if cast_type.as_ref() != left_column.datatype() { - left_cast.push(ScalarExpression::type_cast( - ScalarExpression::column_expr(*left_schema, position), - cast_type.clone(), - arena, - )?); - } else { - left_cast.push(ScalarExpression::column_expr(*left_schema, position)); - } - if cast_type.as_ref() != right_column.datatype() { - right_cast.push(ScalarExpression::type_cast( - ScalarExpression::column_expr(*right_schema, position), - cast_type.clone(), - arena, - )?); - } else { - right_cast.push(ScalarExpression::column_expr(*right_schema, position)); - } + let cast_type = LogicalType::max_logical_type( + arena.column(*left_schema).datatype(), + arena.column(*right_schema).datatype(), + )? + .into_owned(); + + let left_expr = ScalarExpression::column_expr(*left_schema, position) + .type_cast(Cow::Borrowed(&cast_type), arena)?; + left_cast.push(arena.alloc_expression(left_expr)); + + let right_expr = ScalarExpression::column_expr(*right_schema, position) + .type_cast(Cow::Owned(cast_type), arena)?; + right_cast.push(arena.alloc_expression(right_expr)); } if !left_cast.is_empty() { @@ -1036,7 +1012,9 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .iter() .cloned() .enumerate() - .map(|(position, column)| ScalarExpression::column_expr(column, position)) + .map(|(position, column)| { + arena.alloc_expression(ScalarExpression::column_expr(column, position)) + }) .collect_vec(); let union_op = Operator::Union(UnionOperator { @@ -1068,13 +1046,17 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .iter() .cloned() .enumerate() - .map(|(position, column)| ScalarExpression::column_expr(column, position)) + .map(|(position, column)| { + arena.alloc_expression(ScalarExpression::column_expr(column, position)) + }) .collect_vec(); let right_distinct_exprs = right_schema .iter() .cloned() .enumerate() - .map(|(position, column)| ScalarExpression::column_expr(column, position)) + .map(|(position, column)| { + arena.alloc_expression(ScalarExpression::column_expr(column, position)) + }) .collect_vec(); left_plan = self.bind_distinct(left_plan, left_distinct_exprs)?; @@ -1130,18 +1112,15 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' alias_column.set_ref_table(table_alias.clone(), column_id, is_temp); let alias_column = arena.alloc_column(alias_column); - let alias_column_expr = ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr(column, position)), - alias: AliasType::Expr(Box::new(ScalarExpression::column_expr( - alias_column, - position, - ))), - }; - self.context.add_alias( - Some(table_alias.to_string()), - alias, - alias_column_expr.clone(), - ); + let expr = arena.alloc_expression(ScalarExpression::column_expr(column, position)); + let alias_expr = + arena.alloc_expression(ScalarExpression::column_expr(alias_column, position)); + let alias_column_expr = arena.alloc_expression(ScalarExpression::Alias { + expr, + alias: AliasType::Expr(alias_expr), + }); + self.context + .add_alias(Some(table_alias.to_string()), alias, alias_column_expr); alias_exprs.push(alias_column_expr); } self.context.add_table_alias(table_alias, table_name); @@ -1171,13 +1150,13 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' }; let source_column = arena.alloc_column(source_column); - source_exprs.push(ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr(column, position)), - alias: AliasType::Expr(Box::new(ScalarExpression::column_expr( - source_column, - position, - ))), - }); + let expr = arena.alloc_expression(ScalarExpression::column_expr(column, position)); + let alias_expr = + arena.alloc_expression(ScalarExpression::column_expr(source_column, position)); + source_exprs.push(arena.alloc_expression(ScalarExpression::Alias { + expr, + alias: AliasType::Expr(alias_expr), + })); } Self::build_project_plan(plan, source_exprs) @@ -1343,14 +1322,14 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' pub(crate) fn bind_table_column_refs( context: &BinderContext<'a, T>, arena: &mut crate::planner::PlanArena, - exprs: &mut Vec, + exprs: &mut Vec, table_name: TableName, is_qualified_wildcard: bool, ) -> Result<(), DatabaseError> { let (source, position_offset) = Self::resolve_source_columns_in_scope(context, table_name.as_ref())?; - let fn_not_on_using = |column: &ColumnRef| { + let fn_not_on_using = |column: &ColumnRef, arena: &crate::planner::PlanArena<'_>| { let column_catalog = arena.column(*column); if context.using.is_empty() { return Some(&table_name) == column_catalog.table_name(); @@ -1381,13 +1360,13 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' else { continue; }; - if !fn_not_on_using(column) { + if !fn_not_on_using(column, arena) { continue; } - exprs.push(ScalarExpression::column_expr( + exprs.push(arena.alloc_expression(ScalarExpression::column_expr( *column, position_offset + position, - )); + ))); pushed_alias_columns = true; } @@ -1396,13 +1375,13 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } for (position, column) in source.schema().iter().enumerate() { - if !fn_not_on_using(column) { + if !fn_not_on_using(column, arena) { continue; } - exprs.push(ScalarExpression::column_expr( + exprs.push(arena.alloc_expression(ScalarExpression::column_expr( *column, position_offset + position, - )); + ))); } Ok(()) } @@ -1417,14 +1396,11 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' ) -> Result { let left_len = left.output_schema(arena).len(); right.output_schema(arena); - let mut on = self.bind_join_constraint( - join_type, - constraint, - left.output_schema(arena), - right.output_schema(arena), - arena, - )?; - Self::localize_join_condition_from_join_scope(&mut on, left_len)?; + let left_schema = left.output_schema(arena); + let right_schema = right.output_schema(arena); + let mut on = + self.bind_join_constraint(join_type, constraint, &left_schema, &right_schema, arena)?; + Self::localize_join_condition_from_join_scope(&mut on, left_len, arena)?; Ok(LJoinOperator::build( left, @@ -1438,7 +1414,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' pub(crate) fn bind_where_expr( &mut self, mut children: LogicalPlan, - mut predicate: ScalarExpression, + predicate: ExprRef, arena: &mut crate::planner::PlanArena, ) -> Result { self.context.step(QueryBindStep::Where); @@ -1461,7 +1437,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' uses_mark_apply = Some(true); let left_schema = children.output_schema(arena).clone(); let (plan, predicates) = Self::prepare_mark_apply( - &mut predicate, + predicate, &output_column, left_schema.as_ref(), plan, @@ -1493,18 +1469,20 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } uses_mark_apply = Some(true); if correlated { - quantified_predicate = - Self::rewrite_correlated_quantified_predicate(quantified_predicate); + quantified_predicate = Self::rewrite_correlated_quantified_predicate( + quantified_predicate, + arena, + ); } let left_schema = children.output_schema(arena).clone(); let (plan, predicates) = Self::prepare_mark_apply( - &mut predicate, + predicate, &output_column, left_schema.as_ref(), plan, correlated, true, - vec![quantified_predicate], + vec![arena.alloc_expression(quantified_predicate)], arena, )?; children = MarkApplyOperator::build_quantified( @@ -1533,7 +1511,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' children, plan, JoinType::Inner, - std::iter::once(predicate.clone()), + std::iter::once(predicate), true, arena, )?; @@ -1546,7 +1524,9 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .iter() .cloned() .enumerate() - .map(|(position, column)| ScalarExpression::column_expr(column, position)) + .map(|(position, column)| { + arena.alloc_expression(ScalarExpression::column_expr(column, position)) + }) .collect(); let filter = FilterOperator::build(predicate, children, false); return Ok(LogicalPlan::new( @@ -1563,7 +1543,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn ensure_mark_apply_right_outputs( plan: &mut LogicalPlan, - predicates: &[ScalarExpression], + predicates: &[ExprRef], arena: &mut crate::planner::PlanArena, ) -> Result, DatabaseError> { let output_schema = plan.output_schema(arena).clone(); @@ -1593,8 +1573,9 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } } if referenced { - op.exprs - .push(ScalarExpression::column_expr(*column, position)); + op.exprs.push( + arena.alloc_expression(ScalarExpression::column_expr(*column, position)), + ); appended_outputs.push(AppendedRightOutput { column: *column, child_position: position, @@ -1611,22 +1592,21 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' #[allow(clippy::too_many_arguments)] fn prepare_mark_apply( - predicate: &mut ScalarExpression, + mut predicate: ExprRef, output_column: &ColumnRef, left_schema: &Schema, plan: LogicalPlan, correlated: bool, preserve_projection: bool, - mut apply_predicates: Vec, + mut apply_predicates: Vec, arena: &mut crate::planner::PlanArena, - ) -> Result<(LogicalPlan, Vec), DatabaseError> { + ) -> Result<(LogicalPlan, Vec), DatabaseError> { let left_len = left_schema.len(); MarkerPositionGlobalizer { output_column, left_len, - arena, } - .visit(predicate)?; + .visit(&mut predicate, arena)?; let (mut plan, correlated_filters) = if correlated { Self::prepare_correlated_subquery_plan(plan, left_schema, preserve_projection, arena)? @@ -1647,25 +1627,27 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } } let right_schema = plan.output_schema(arena); - for expr in apply_predicates.iter_mut() { + for expr in &mut apply_predicates { RightSidePositionGlobalizer { - right_schema, + right_schema: &right_schema, left_len, - arena, } - .visit(expr)?; + .visit(expr, arena)?; } Ok((plan, apply_predicates)) } - fn rewrite_correlated_quantified_predicate(predicate: ScalarExpression) -> ScalarExpression { - let strip_projection_alias = |expr: Box| match *expr { + fn rewrite_correlated_quantified_predicate( + predicate: ScalarExpression, + arena: &PlanArena<'_>, + ) -> ScalarExpression { + let strip_projection_alias = |expr| match arena.expression(expr) { ScalarExpression::Alias { expr, alias: AliasType::Expr(_), - } => expr, - expr => Box::new(expr), + } => *expr, + _ => expr, }; match predicate { @@ -1716,7 +1698,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } fn expr_has_correlated_refs( - expr: &ScalarExpression, + expr: ExprRef, left_schema: &Schema, arena: &mut crate::planner::PlanArena, ) -> Result { @@ -1727,31 +1709,32 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' }) } - fn split_conjuncts(expr: ScalarExpression, exprs: &mut Vec) { - match expr.unpack_alias() { + fn split_conjuncts(expr: ExprRef, exprs: &mut Vec, arena: &PlanArena<'_>) { + let expr = expr.unpack_alias(arena); + match arena.expression(expr) { ScalarExpression::Binary { op: BinaryOperator::And, left_expr, right_expr, .. } => { - Self::split_conjuncts(*left_expr, exprs); - Self::split_conjuncts(*right_expr, exprs); + Self::split_conjuncts(*left_expr, exprs, arena); + Self::split_conjuncts(*right_expr, exprs, arena); } - expr => exprs.push(expr), + _ => exprs.push(expr), } } - fn combine_conjuncts(exprs: Vec) -> Option { - exprs - .into_iter() - .reduce(|acc, expr| ScalarExpression::Binary { + fn combine_conjuncts(exprs: Vec, arena: &mut PlanArena<'_>) -> Option { + exprs.into_iter().reduce(|acc, expr| { + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(acc), - right_expr: Box::new(expr), + left_expr: acc, + right_expr: expr, evaluator: None, ty: LogicalType::Boolean, }) + }) } fn prepare_correlated_subquery_plan( @@ -1759,7 +1742,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' left_schema: &Schema, preserve_projection: bool, arena: &mut crate::planner::PlanArena, - ) -> Result<(LogicalPlan, Vec), DatabaseError> { + ) -> Result<(LogicalPlan, Vec), DatabaseError> { match plan.childrens.as_ref() { Childrens::Only(_) => {} Childrens::Twins { .. } => { @@ -1788,15 +1771,15 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' )?; let mut local_filters = Vec::new(); let mut predicates = Vec::new(); - Self::split_conjuncts(op.predicate, &mut predicates); + Self::split_conjuncts(op.predicate, &mut predicates, arena); for predicate in predicates { - if Self::expr_has_correlated_refs(&predicate, left_schema, arena)? { + if Self::expr_has_correlated_refs(predicate, left_schema, arena)? { correlated_filters.push(predicate); } else { local_filters.push(predicate); } } - let plan = if let Some(predicate) = Self::combine_conjuncts(local_filters) { + let plan = if let Some(predicate) = Self::combine_conjuncts(local_filters, arena) { FilterOperator::build(predicate, child, op.having) } else { child @@ -1819,9 +1802,9 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' if !preserve_projection || Self::is_temp_alias_projection(&op.exprs, arena) { Ok((child, correlated_filters)) } else { - let mut binder = ProjectionOutputBinder::new(&op.exprs, arena); - for expr in correlated_filters.iter_mut() { - binder.visit(expr)?; + let mut binder = ProjectionOutputBinder::new(&op.exprs); + for expr in &mut correlated_filters { + binder.visit(expr, arena)?; } Ok(( LogicalPlan::new(Operator::Project(op), Childrens::Only(Box::new(child))), @@ -1865,19 +1848,19 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn bind_having( &mut self, children: LogicalPlan, - mut having: ScalarExpression, + mut having: ExprRef, arena: &mut crate::planner::PlanArena, ) -> Result { self.context.step(QueryBindStep::Having); - self.validate_having_orderby(&having)?; + self.validate_having_orderby(having, arena)?; self.bind_aggregate_output_exprs(std::iter::once(&mut having), arena)?; Ok(FilterOperator::build(having, children, true)) } pub(crate) fn build_project_plan( children: LogicalPlan, - select_list: Vec, + select_list: Vec, ) -> LogicalPlan { LogicalPlan::new( Operator::Project(ProjectOperator { exprs: select_list }), @@ -1888,7 +1871,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' pub(crate) fn bind_project( &mut self, mut children: LogicalPlan, - mut select_list: Vec, + mut select_list: Vec, arena: &mut crate::planner::PlanArena, ) -> Result { self.context.step(QueryBindStep::Project); @@ -1913,13 +1896,12 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' let left_len = children.output_schema(arena).len(); let right_schema = plan.output_schema(arena); - for expr in select_list.iter_mut() { + for expr in &mut select_list { RightSidePositionGlobalizer { right_schema, left_len, - arena, } - .visit(expr)?; + .visit(expr, arena)?; } children = ScalarApplyOperator::build(children, plan); @@ -1956,7 +1938,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' pub fn extract_select_join( &mut self, - select_items: &mut [ScalarExpression], + select_items: &mut [ExprRef], arena: &mut crate::planner::PlanArena, ) { if self.context.bind_table.len() < 2 { @@ -1985,8 +1967,10 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' table_force_nullable.push((table_name, table, left_table_force_nullable)); } - for column in select_items { - if let ScalarExpression::ColumnRef { column, .. } = column { + for expr in select_items { + let mut expression = + std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + if let ScalarExpression::ColumnRef { column, .. } = &mut expression { let _ = table_force_nullable .iter() .find(|(table_name, _source, _)| { @@ -2001,6 +1985,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' } }); } + *arena.expression_mut(*expr) = expression; } } @@ -2015,7 +2000,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' match constraint { JoinConstraintInput::On(expr) => { // left and right columns that match equi-join pattern - let mut on_keys: Vec<(ScalarExpression, ScalarExpression)> = vec![]; + let mut on_keys: Vec<(ExprRef, ExprRef)> = vec![]; // expression that didn't match equi-join pattern let mut filter = vec![]; @@ -2029,15 +2014,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' )?; // combine multiple filter exprs into one BinaryExpr - let join_filter = filter - .into_iter() - .reduce(|acc, expr| ScalarExpression::Binary { - op: BinaryOperator::And, - left_expr: Box::new(acc), - right_expr: Box::new(expr), - evaluator: None, - ty: LogicalType::Boolean, - }); + let join_filter = Self::combine_conjuncts(filter, arena); Ok(JoinCondition::On { on: on_keys, filter: join_filter, @@ -2055,7 +2032,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .find(|(_, column)| arena.column(**column).name() == name) } - let mut on_keys: Vec<(ScalarExpression, ScalarExpression)> = Vec::new(); + let mut on_keys: Vec<(ExprRef, ExprRef)> = Vec::new(); for name in names { let (Some((left_position, left_column)), Some((right_position, right_column))) = ( @@ -2074,13 +2051,15 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' right_column, left_schema.len() + right_position, )?; - on_keys.push(( - ScalarExpression::column_expr(*left_column, left_position), - ScalarExpression::column_expr( - *right_column, - left_schema.len() + right_position, - ), + let left_expr = arena.alloc_expression(ScalarExpression::column_expr( + *left_column, + left_position, )); + let right_expr = arena.alloc_expression(ScalarExpression::column_expr( + *right_column, + left_schema.len() + right_position, + )); + on_keys.push((left_expr, right_expr)); } Ok(JoinCondition::On { on: on_keys, @@ -2095,7 +2074,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .map(|column| arena.column(*column).name().to_string()) .collect() }; - let mut on_keys: Vec<(ScalarExpression, ScalarExpression)> = Vec::new(); + let mut on_keys: Vec<(ExprRef, ExprRef)> = Vec::new(); for name in fn_names(left_schema).intersection(&fn_names(right_schema)) { if let ( @@ -2111,11 +2090,14 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' .enumerate() .find(|(_, column)| arena.column(**column).name() == name), ) { - let left_expr = ScalarExpression::column_expr(*left_column, left_position); - let right_expr = ScalarExpression::column_expr( + let left_expr = arena.alloc_expression(ScalarExpression::column_expr( + *left_column, + left_position, + )); + let right_expr = arena.alloc_expression(ScalarExpression::column_expr( *right_column, left_schema.len() + right_position, - ); + )); self.context.add_using( name.clone(), @@ -2148,9 +2130,9 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' /// foo = bar AND baz > 1 => accum=[(foo, bar)] accum_filter=[baz > 1] /// ``` fn extract_join_keys( - expr: ScalarExpression, - accum: &mut Vec<(ScalarExpression, ScalarExpression)>, - accum_filter: &mut Vec, + expr: ExprRef, + accum: &mut Vec<(ExprRef, ExprRef)>, + accum_filter: &mut Vec, left_schema: &Schema, right_schema: &Schema, arena: &crate::planner::PlanArena, @@ -2165,17 +2147,20 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn_contains(left_schema, column) || fn_contains(right_schema, column) }; - match expr.unpack_alias() { + let expr = expr.unpack_alias(arena); + match arena.expression(expr) { ScalarExpression::Binary { left_expr, right_expr, op, - ty, .. } => { match op { BinaryOperator::Eq => { - match (left_expr.unpack_alias_ref(), right_expr.unpack_alias_ref()) { + match ( + left_expr.unpack_alias_ref(arena), + right_expr.unpack_alias_ref(arena), + ) { // example: foo = bar ( ScalarExpression::ColumnRef { column: l, .. }, @@ -2189,25 +2174,13 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' { accum.push((*right_expr, *left_expr)); } else if fn_or_contains(*l) || fn_or_contains(*r) { - accum_filter.push(ScalarExpression::Binary { - left_expr, - right_expr, - op, - ty, - evaluator: None, - }); + accum_filter.push(expr); } } (ScalarExpression::ColumnRef { column, .. }, _) | (_, ScalarExpression::ColumnRef { column, .. }) => { if fn_or_contains(*column) { - accum_filter.push(ScalarExpression::Binary { - left_expr, - right_expr, - op, - ty, - evaluator: None, - }); + accum_filter.push(expr); } } _other => { @@ -2219,13 +2192,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn_or_contains(*column) })? { - accum_filter.push(ScalarExpression::Binary { - left_expr, - right_expr, - op, - ty, - evaluator: None, - }); + accum_filter.push(expr); } } } @@ -2250,13 +2217,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' )?; } BinaryOperator::Or => { - accum_filter.push(ScalarExpression::Binary { - left_expr, - right_expr, - op, - ty, - evaluator: None, - }); + accum_filter.push(expr); } _ => { if left_expr @@ -2265,18 +2226,12 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' fn_or_contains(*column) })? { - accum_filter.push(ScalarExpression::Binary { - left_expr, - right_expr, - op, - ty, - evaluator: None, - }); + accum_filter.push(expr); } } } } - expr => { + _ => { if expr.all_referenced_columns(arena, |_, column| fn_or_contains(*column))? { // example: baz > 1 accum_filter.push(expr); @@ -2301,18 +2256,16 @@ mod tests { MarkApplyKind, MarkApplyOperator, MarkApplyQuantifier, }; use crate::planner::operator::Operator; - use crate::planner::{Childrens, LogicalPlan, PlanArena}; + use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::LogicalType; - fn test_column(arena: &mut PlanArena, name: &str, position: usize) -> ScalarExpression { - ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new( - name.to_string(), - true, - ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(), - )), - position, - ) + fn test_column(arena: &mut PlanArena, name: &str, position: usize) -> ExprRef { + let column = arena.alloc_column(ColumnCatalog::new( + name.to_string(), + true, + ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(), + )); + arena.alloc_expression(ScalarExpression::column_expr(column, position)) } #[test] @@ -2361,40 +2314,41 @@ mod tests { ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(), )); let right_schema = vec![right_column]; - let mut expr = ScalarExpression::Binary { + let left_expr = arena.alloc_expression(ScalarExpression::column_expr(left_column, 0)); + let right_expr = arena.alloc_expression(ScalarExpression::column_expr(right_column, 0)); + let mut expr = arena.alloc_expression(ScalarExpression::Binary { op: crate::expression::BinaryOperator::Eq, - left_expr: Box::new(ScalarExpression::column_expr(left_column, 0)), - right_expr: Box::new(ScalarExpression::column_expr(right_column, 0)), + left_expr, + right_expr, evaluator: None, ty: LogicalType::Boolean, - }; + }); RightSidePositionGlobalizer { right_schema: &right_schema, left_len: 2, - arena: &arena, } - .visit(&mut expr)?; + .visit(&mut expr, &mut arena)?; let ScalarExpression::Binary { left_expr, right_expr, .. - } = expr + } = arena.expression(expr) else { unreachable!() }; let ScalarExpression::ColumnRef { position: left_position, .. - } = left_expr.as_ref() + } = arena.expression(*left_expr) else { unreachable!() }; let ScalarExpression::ColumnRef { position: right_position, .. - } = right_expr.as_ref() + } = arena.expression(*right_expr) else { unreachable!() }; @@ -2407,21 +2361,23 @@ mod tests { fn test_projection_output_binder_rewrites_to_project_slot() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); - let project_output = ScalarExpression::Alias { - expr: Box::new(test_column(&mut arena, "c1", 0)), + let project_inner = test_column(&mut arena, "c1", 0); + let project_output = arena.alloc_expression(ScalarExpression::Alias { + expr: project_inner, alias: AliasType::Name("v".to_string()), - }; - let mut expr = ScalarExpression::Alias { - expr: Box::new(test_column(&mut arena, "c1", 0)), + }); + let expr_inner = test_column(&mut arena, "c1", 0); + let mut expr = arena.alloc_expression(ScalarExpression::Alias { + expr: expr_inner, alias: AliasType::Name("v".to_string()), - }; + }); - ProjectionOutputBinder::new(std::slice::from_ref(&project_output), &mut arena) - .visit(&mut expr)?; + ProjectionOutputBinder::new(std::slice::from_ref(&project_output)) + .visit(&mut expr, &mut arena)?; - let expected = - ScalarExpression::column_expr(project_output.output_column_ref(&mut arena), 0); - assert!(expr.eq_ignore_colref_pos(&expected, &arena)); + let output_column = project_output.output_column_ref(&mut arena); + let expected = arena.alloc_expression(ScalarExpression::column_expr(output_column, 0)); + assert!(expr.eq_ignore_colref_pos(expected, &arena)); Ok(()) } @@ -2555,16 +2511,16 @@ mod tests { } } - fn collect_column_positions(expr: &ScalarExpression, positions: &mut Vec) { - match expr.unpack_alias_ref() { + fn collect_column_positions(expr: ExprRef, arena: &PlanArena, positions: &mut Vec) { + match arena.expression(expr.unpack_alias(arena)) { ScalarExpression::ColumnRef { position, .. } => positions.push(*position), ScalarExpression::Binary { left_expr, right_expr, .. } => { - collect_column_positions(left_expr, positions); - collect_column_positions(right_expr, positions); + collect_column_positions(*left_expr, arena, positions); + collect_column_positions(*right_expr, arena, positions); } _ => {} } @@ -2573,8 +2529,11 @@ mod tests { #[test] fn test_multiple_scalar_subqueries_in_where_rebind_positions() -> Result<(), DatabaseError> { let table_states = build_t1_table()?; - let plan = - table_states.plan("select * from t1 where c1 <= (select 4) and c1 > (select 1)")?; + let mut arena = PlanArena::new(&table_states.table_arena); + let plan = table_states.plan_with_arena( + "select * from t1 where c1 <= (select 4) and c1 > (select 1)", + &mut arena, + )?; let outer_join = find_top_join(&plan).expect("expected scalar subqueries to introduce a join"); let Operator::Join(op) = &outer_join.operator else { @@ -2590,12 +2549,11 @@ mod tests { else { panic!("expected join filter") }; - let mut arena = PlanArena::new(&table_states.table_arena); let mut left_plan = left.as_ref().clone(); let left_len = left_plan.output_schema(&mut arena).len(); let mut positions = Vec::new(); - collect_column_positions(filter, &mut positions); + collect_column_positions(*filter, &arena, &mut positions); assert_eq!(positions, vec![0, left_len - 1, 0, left_len]); diff --git a/src/binder/update.rs b/src/binder/update.rs index e532d110..9e1270d9 100644 --- a/src/binder/update.rs +++ b/src/binder/update.rs @@ -15,10 +15,9 @@ use crate::binder::Binder; use crate::catalog::{ColumnRef, TableName}; use crate::errors::DatabaseError; -use crate::expression::ScalarExpression; use crate::planner::operator::update::UpdateOperator; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -26,7 +25,7 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_update( &mut self, table_name: TableName, - value_exprs: Vec<(ColumnRef, ScalarExpression)>, + value_exprs: Vec<(ColumnRef, ExprRef)>, input: LogicalPlan, ) -> Result { Ok(LogicalPlan::new( diff --git a/src/binder/window.rs b/src/binder/window.rs index 9627016f..fd33e0b6 100644 --- a/src/binder/window.rs +++ b/src/binder/window.rs @@ -22,42 +22,46 @@ use crate::planner::operator::sort::SortField; use crate::planner::operator::sort::SortOperator; use crate::planner::operator::window::WindowOperator; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan, PlanArena}; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::value::DataValue; use crate::types::LogicalType; -struct WindowCollector<'a, 'p> { - arena: &'a mut PlanArena<'p>, +struct WindowCollector { windows: Vec<(WindowCall, ColumnRef)>, } -impl ExprVisitorMut<'_> for WindowCollector<'_, '_> { - fn visit(&mut self, expr: &mut ScalarExpression) -> Result<(), DatabaseError> { - let ScalarExpression::WindowCall(window) = expr else { - return walk_mut_expr(self, expr); +impl ExprVisitorMut for WindowCollector { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + let ScalarExpression::WindowCall(window) = arena.expression(*expr) else { + return walk_mut_expr(self, expr, arena); }; if let Some((_, output_column)) = self .windows .iter() .find(|(candidate, _)| candidate == window) { - *expr = ScalarExpression::column_expr(*output_column, 0); + *arena.expression_mut(*expr) = ScalarExpression::column_expr(*output_column, 0); return Ok(()); } - let output_name = expr.output_name(self.arena); - let ScalarExpression::WindowCall(window) = std::mem::replace(expr, ScalarExpression::Empty) + let output_name = expr.output_name(arena); + let ScalarExpression::WindowCall(window) = + std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty) else { unreachable!() }; - let output_column = self.arena.alloc_column(ColumnCatalog::new( + let output_column = arena.alloc_column(ColumnCatalog::new( output_name, true, ColumnDesc::new(window.function.ty.clone(), None, false, None)?, )); self.windows.push((window, output_column)); - *expr = ScalarExpression::column_expr(output_column, 0); + *arena.expression_mut(*expr) = ScalarExpression::column_expr(output_column, 0); Ok(()) } } @@ -67,11 +71,12 @@ struct WindowOutputBinder<'a> { base_position: usize, } -impl ExprVisitorMut<'_> for WindowOutputBinder<'_> { +impl ExprVisitorMut for WindowOutputBinder<'_> { fn visit_column_ref( &mut self, column: &mut ColumnRef, position: &mut usize, + _arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if let Some(output_position) = self .groups @@ -86,7 +91,7 @@ impl ExprVisitorMut<'_> for WindowOutputBinder<'_> { } struct WindowGroup { - partition_by: Vec, + partition_by: Vec, order_by: Vec, functions: Vec, output_columns: Vec, @@ -96,8 +101,8 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_window_function( &mut self, kind: WindowFunctionKind, - args: Vec, - partition_by: Vec, + args: Vec, + partition_by: Vec, order_by: Vec, arena: &mut PlanArena, ) -> Result { @@ -114,7 +119,7 @@ impl> Binder<'_, '_, T, A> .chain(&partition_by) .chain(order_by.iter().map(|field| &field.expr)) { - if expr.has_window_call()? { + if expr.has_window_call(arena)? { return Err(DatabaseError::UnsupportedStmt( "window functions cannot be nested".to_string(), )); @@ -155,20 +160,19 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_window( &mut self, mut children: LogicalPlan, - select_list: &mut [ScalarExpression], + select_list: &mut [ExprRef], order_by: &mut Option>, arena: &mut PlanArena, ) -> Result { let mut collector = WindowCollector { - arena, windows: Vec::new(), }; for expr in select_list.iter_mut() { - collector.visit(expr)?; + collector.visit(expr, arena)?; } if let Some(order_by) = order_by.as_mut() { for field in order_by { - collector.visit(&mut field.expr)?; + collector.visit(&mut field.expr, arena)?; } } if collector.windows.is_empty() { @@ -205,7 +209,7 @@ impl> Binder<'_, '_, T, A> .iter_mut() .chain(order_by.iter_mut().flatten().map(|field| &mut field.expr)) { - output_binder.visit(expr)?; + output_binder.visit(expr, arena)?; } for group in groups { diff --git a/src/catalog/column.rs b/src/catalog/column.rs index 03fc60f1..67741ca3 100644 --- a/src/catalog/column.rs +++ b/src/catalog/column.rs @@ -14,7 +14,7 @@ use crate::catalog::TableName; use crate::errors::DatabaseError; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -172,11 +172,14 @@ impl ColumnCatalog { &self.desc.column_datatype } - pub(crate) fn default_value(&self) -> Result, DatabaseError> { + pub(crate) fn default_value( + &self, + arena: &PlanArena<'_>, + ) -> Result, DatabaseError> { self.desc .default .as_ref() - .map(|expr| expr.eval::<&Tuple>(None)) + .map(|expr| arena.expression(*expr).eval::<&Tuple>(arena, None)) .transpose() } @@ -231,7 +234,7 @@ pub struct ColumnDesc { pub(crate) column_datatype: LogicalType, primary: Option, is_unique: bool, - pub(crate) default: Option, + pub(crate) default: Option, } impl ColumnDesc { @@ -239,16 +242,8 @@ impl ColumnDesc { column_datatype: LogicalType, primary: Option, is_unique: bool, - default: Option, + default: Option, ) -> Result { - if let Some(expr) = &default { - let table_arena = crate::planner::TableArenaCell::default(); - let plan_arena = crate::planner::PlanArena::new(&table_arena); - if expr.has_table_ref_column(&plan_arena)? { - return Err(DatabaseError::DefaultNotColumnRef); - } - } - Ok(ColumnDesc { column_datatype, primary, diff --git a/src/db.rs b/src/db.rs index ec9e2bad..249f89eb 100644 --- a/src/db.rs +++ b/src/db.rs @@ -1459,7 +1459,7 @@ pub(crate) mod test { ); let mut source_plan_arena = PlanArena::new(kite_sql.state.table_arena()); let source_plan = binder.bind(&stmt, &mut source_plan_arena)?; - let (best_plan, _best_plan_arena) = + let (best_plan, best_plan_arena) = kite_sql .state .build_plan([], &transaction, |binder, arena| binder.bind(&stmt, arena))?; @@ -1480,14 +1480,14 @@ pub(crate) mod test { let ScalarExpression::ColumnRef { position: left_position, .. - } = on[0].0.unpack_alias_ref() + } = source_plan_arena.expression(on[0].0.unpack_alias(&source_plan_arena)) else { unreachable!("expected left join key column ref"); }; let ScalarExpression::ColumnRef { position: right_position, .. - } = on[0].1.unpack_alias_ref() + } = source_plan_arena.expression(on[0].1.unpack_alias(&source_plan_arena)) else { unreachable!("expected right join key column ref"); }; @@ -1508,7 +1508,7 @@ pub(crate) mod test { let ScalarExpression::ColumnRef { position: right_position, .. - } = on[0].1.unpack_alias_ref() + } = best_plan_arena.expression(on[0].1.unpack_alias(&best_plan_arena)) else { unreachable!("expected right join key column ref"); }; @@ -1542,7 +1542,7 @@ pub(crate) mod test { ); let mut source_plan_arena = PlanArena::new(kite_sql.state.table_arena()); let source_plan = binder.bind(&stmt, &mut source_plan_arena)?; - let (best_plan, _best_plan_arena) = + let (best_plan, best_plan_arena) = kite_sql .state .build_plan([], &transaction, |binder, arena| binder.bind(&stmt, arena))?; @@ -1562,14 +1562,14 @@ pub(crate) mod test { let ScalarExpression::ColumnRef { position: left_position, .. - } = on[0].0.unpack_alias_ref() + } = source_plan_arena.expression(on[0].0.unpack_alias(&source_plan_arena)) else { unreachable!("expected left join key column ref"); }; let ScalarExpression::ColumnRef { position: right_position, .. - } = on[0].1.unpack_alias_ref() + } = source_plan_arena.expression(on[0].1.unpack_alias(&source_plan_arena)) else { unreachable!("expected right join key column ref"); }; @@ -1579,7 +1579,7 @@ pub(crate) mod test { unreachable!("expected join filter"); }; let mut referenced_columns = Vec::new(); - filter.visit_referenced_columns(&mut source_plan_arena, &mut |_, column| { + filter.all_referenced_columns(&source_plan_arena, |_, column| { referenced_columns.push(*column); true })?; @@ -1602,14 +1602,14 @@ pub(crate) mod test { let ScalarExpression::ColumnRef { position: left_position, .. - } = on[0].0.unpack_alias_ref() + } = best_plan_arena.expression(on[0].0.unpack_alias(&best_plan_arena)) else { unreachable!("expected left join key column ref"); }; let ScalarExpression::ColumnRef { position: right_position, .. - } = on[0].1.unpack_alias_ref() + } = best_plan_arena.expression(on[0].1.unpack_alias(&best_plan_arena)) else { unreachable!("expected right join key column ref"); }; @@ -1639,7 +1639,7 @@ pub(crate) mod test { "SELECT o.x, t.y FROM onecolumn o INNER JOIN twocolumn t ON (o.x=t.x AND t.y=53)", )?; let transaction = kite_sql.storage.transaction()?; - let (best_plan, _best_plan_arena) = + let (best_plan, best_plan_arena) = kite_sql .state .build_plan([], &transaction, |binder, arena| binder.bind(&stmt, arena))?; @@ -1659,14 +1659,14 @@ pub(crate) mod test { let ScalarExpression::ColumnRef { position: left_position, .. - } = on[0].0.unpack_alias_ref() + } = best_plan_arena.expression(on[0].0.unpack_alias(&best_plan_arena)) else { unreachable!("expected left join key column ref"); }; let ScalarExpression::ColumnRef { position: right_position, .. - } = on[0].1.unpack_alias_ref() + } = best_plan_arena.expression(on[0].1.unpack_alias(&best_plan_arena)) else { unreachable!("expected right join key column ref"); }; @@ -1680,20 +1680,20 @@ pub(crate) mod test { left_expr, right_expr, .. - } = filter_op.predicate + } = best_plan_arena.expression(filter_op.predicate) else { unreachable!("expected binary filter predicate"); }; let ScalarExpression::ColumnRef { position: filter_position, .. - } = left_expr.unpack_alias_ref() + } = best_plan_arena.expression(left_expr.unpack_alias(&best_plan_arena)) else { unreachable!("expected filter column ref"); }; assert_eq!(*filter_position, 1); assert!(matches!( - *right_expr, + best_plan_arena.expression(*right_expr), ScalarExpression::Constant(DataValue::Int32(53)) )); @@ -1718,10 +1718,10 @@ pub(crate) mod test { let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); - assert!(plan.contains("Projection")); - assert!(plan.contains("Filter (")); - assert!(plan.contains(" > 0")); - assert!(plan.contains("TableScan t1 -> [#")); + assert_eq!( + plan, + "Projection [t1.a, t1.b] [Project => (Sort Option: Follow)] Filter (t1.b > 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.a, t1.b] [SeqScan => (Sort Option: None)]" + ); } // Aggregate { @@ -1740,11 +1740,10 @@ pub(crate) mod test { )?; let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); - assert!(plan.contains("Projection")); - assert!(plan.contains("Aggregate")); - assert!(plan.contains("Filter (")); - assert!(plan.contains(" > 1")); - assert!(plan.contains("TableScan t1 -> [#")); + assert_eq!( + plan, + "Projection [(t1.a + 0), Max((t1.b + 0))] [Project => (Sort Option: Follow)] Aggregate [Max((t1.b + 0))] -> Group By [(t1.a + 0)] [HashAggregate => (Sort Option: None)] Filter (t1.b > 1), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.a, t1.b] [SeqScan => (Sort Option: None)]" + ); } { let statement = crate::db::prepare("explain select *, $1 from (select * from t1 where b > $2) left join (select * from t1 where a > $3) on a > $4")?; @@ -1760,238 +1759,11 @@ pub(crate) mod test { )?; let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); - assert!(plan.contains("Projection")); - assert!(plan.contains("LeftOuter Join")); - assert!(plan.contains("9")); - assert!(plan.contains("0")); - assert!(plan.contains("1")); - assert!(plan.contains("TableScan t1 -> [#")); - } - - Ok(()) - } - - // FIXME: keep this as a unit test instead of SLT for now. The current - // sqllogictest runner does not reliably match the pretty-printed multi-line - // EXPLAIN output produced by correlated IN, even though the plan itself is stable. - #[test] - fn test_subquery_explain_uses_parameterized_index_for_in() -> Result<(), DatabaseError> { - let temp_dir = TempDir::new().expect("unable to create temporary working directory"); - let mut kite_sql = DataBaseBuilder::path(temp_dir.path()).build_rocksdb()?; - - kite_sql.ddl("create table in_outer(id int primary key, a int)")?; - kite_sql.ddl("create table in_inner(id int primary key, v int)")?; - kite_sql.ddl("create table in_inner_nn(id int primary key, v int)")?; - kite_sql.ddl("create index in_inner_v_index on in_inner(v)")?; - kite_sql.ddl("create index in_inner_nn_v_index on in_inner_nn(v)")?; - - kite_sql - .run("insert into in_outer values (0, null), (1, 1), (2, 2), (3, 3)")? - .done()?; - kite_sql - .run("insert into in_inner values (0, 2), (1, null)")? - .done()?; - kite_sql - .run("insert into in_inner_nn values (0, 2)")? - .done()?; - - kite_sql.ddl("create table in_outer_flag(id int primary key, a int, b int)")?; - kite_sql.ddl("create table in_inner_flag(id int primary key, v int, flag int)")?; - kite_sql.ddl("create table in_inner_flag_nn(id int primary key, v int, flag int)")?; - kite_sql.ddl("create index in_inner_flag_v_index on in_inner_flag(v)")?; - kite_sql.ddl("create index in_inner_flag_nn_v_index on in_inner_flag_nn(v)")?; - - kite_sql - .run("insert into in_outer_flag values (0, null, 1), (1, 1, 1), (2, 2, 1), (3, 3, 1)")? - .done()?; - kite_sql - .run("insert into in_inner_flag values (0, 2, 1), (1, null, 1)")? - .done()?; - kite_sql - .run("insert into in_inner_flag_nn values (0, 2, 1)")? - .done()?; - - let collect_plan = |sql: &str| -> Result { - let mut iter = kite_sql.run(sql)?; - let mut lines = Vec::new(); - while let Some(row) = next_tuple_owned(&mut iter)? { - if let Some(DataValue::Utf8 { value, .. }) = row.values.first() { - lines.push(value.clone()); - } - } - iter.done()?; - Ok(lines.join("\n")) - }; - let collect_ids = |sql: &str| -> Result, DatabaseError> { - let mut iter = kite_sql.run(sql)?; - let mut ids = Vec::new(); - while let Some(row) = next_tuple_owned(&mut iter)? { - ids.push(row.values[0].i32().unwrap()); - } - iter.done()?; - Ok(ids) - }; - - let assert_mark_in_uses_parameterized_index = |sql: &str| -> Result<(), DatabaseError> { - let explain_plan = collect_plan(sql)?; - assert!( - explain_plan.contains("MarkAnyApply"), - "unexpected explain plan: {explain_plan}" - ); - assert!( - explain_plan.contains("IndexScan By #") && explain_plan.contains("=> Probe"), - "unexpected explain plan: {explain_plan}" - ); - Ok(()) - }; - - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer where a in (select v from in_inner where in_inner.v = in_outer.a)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer where a not in (select v from in_inner where in_inner.v = in_outer.a)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer where a in (select v from in_inner_nn where in_inner_nn.v = in_outer.a)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer where a not in (select v from in_inner_nn where in_inner_nn.v = in_outer.a)", - )?; - - assert_eq!( - collect_ids( - "select id from in_outer where a in (select v from in_inner where in_inner.v = in_outer.a) order by id", - )?, - vec![2] - ); - assert_eq!( - collect_ids( - "select id from in_outer where a not in (select v from in_inner where in_inner.v = in_outer.a) order by id", - )?, - vec![0, 1, 3] - ); - assert_eq!( - collect_ids( - "select id from in_outer where a in (select v from in_inner_nn where in_inner_nn.v = in_outer.a) order by id", - )?, - vec![2] - ); - assert_eq!( - collect_ids( - "select id from in_outer where a not in (select v from in_inner_nn where in_inner_nn.v = in_outer.a) order by id", - )?, - vec![0, 1, 3] - ); - - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer_flag where a in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer_flag where a not in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer_flag where a in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b)", - )?; - assert_mark_in_uses_parameterized_index( - "explain select id from in_outer_flag where a not in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b)", - )?; - - assert_eq!( - collect_ids( - "select id from in_outer_flag where a in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b) order by id", - )?, - vec![2] - ); - assert_eq!( - collect_ids( - "select id from in_outer_flag where a not in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b) order by id", - )?, - Vec::::new() - ); - assert_eq!( - collect_ids( - "select id from in_outer_flag where a in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b) order by id", - )?, - vec![2] - ); - assert_eq!( - collect_ids( - "select id from in_outer_flag where a not in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b) order by id", - )?, - vec![1, 3] - ); - - Ok(()) - } - - #[test] - fn test_subquery_explain_uses_parameterized_index_for_exists() -> Result<(), DatabaseError> { - let temp_dir = TempDir::new().expect("unable to create temporary working directory"); - let mut kite_sql = DataBaseBuilder::path(temp_dir.path()).build_rocksdb()?; - - kite_sql.ddl("create table exists_outer(id int primary key, a int, b int)")?; - kite_sql.ddl("create table exists_inner(id int primary key, v int, flag int)")?; - kite_sql.ddl("create index exists_inner_v_index on exists_inner(v)")?; - - kite_sql - .run("insert into exists_outer values (0, 1, 1), (1, 1, 2), (2, 2, null), (3, 3, 1)")? - .done()?; - kite_sql - .run("insert into exists_inner values (0, 1, 1), (1, 1, null), (2, 2, 1)")? - .done()?; - - let collect_plan = |sql: &str| -> Result { - let mut iter = kite_sql.run(sql)?; - let mut lines = Vec::new(); - while let Some(row) = next_tuple_owned(&mut iter)? { - if let Some(DataValue::Utf8 { value, .. }) = row.values.first() { - lines.push(value.clone()); - } - } - iter.done()?; - Ok(lines.join("\n")) - }; - let collect_ids = |sql: &str| -> Result, DatabaseError> { - let mut iter = kite_sql.run(sql)?; - let mut ids = Vec::new(); - while let Some(row) = next_tuple_owned(&mut iter)? { - ids.push(row.values[0].i32().unwrap()); - } - iter.done()?; - Ok(ids) - }; - let assert_mark_exists_uses_parameterized_index = |sql: &str| -> Result<(), DatabaseError> { - let explain_plan = collect_plan(sql)?; - assert!( - explain_plan.contains("MarkExistsApply"), - "unexpected explain plan: {explain_plan}" + assert_eq!( + plan, + "Projection [t1.a, t1.b, t1.a, t1.b, 9] [Project => (Sort Option: Follow)] LeftOuter Join Where (t1.a > 0) [NestLoopJoin => (Sort Option: None)] Projection [t1.a, t1.b] [Project => (Sort Option: Follow)] Filter (t1.b > 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.a, t1.b] [SeqScan => (Sort Option: None)] Projection [t1.a, t1.b] [Project => (Sort Option: Follow)] Filter (t1.a > 1), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.a, t1.b] [SeqScan => (Sort Option: None)]" ); - assert!( - explain_plan.contains("IndexScan By #") && explain_plan.contains("=> Probe"), - "unexpected explain plan: {explain_plan}" - ); - Ok(()) - }; - - assert_mark_exists_uses_parameterized_index( - "explain select id from exists_outer where exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b)", - )?; - assert_mark_exists_uses_parameterized_index( - "explain select id from exists_outer where not exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b)", - )?; - - assert_eq!( - collect_ids( - "select id from exists_outer where exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b) order by id", - )?, - vec![0] - ); - assert_eq!( - collect_ids( - "select id from exists_outer where not exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b) order by id", - )?, - vec![1, 2, 3] - ); + } Ok(()) } diff --git a/src/execution/ddl/add_column.rs b/src/execution/ddl/add_column.rs index a1364e91..fc77db96 100644 --- a/src/execution/ddl/add_column.rs +++ b/src/execution/ddl/add_column.rs @@ -86,7 +86,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for AddColumn { return Err(DatabaseError::DuplicateColumn(column.name().to_string())); } - let default_value = column.default_value()?; + let default_value = column.default_value(plan_arena)?; let (unique_index_id, apply) = { let (transaction, table_codec) = arena.transaction_codec_mut(); diff --git a/src/execution/ddl/create_index.rs b/src/execution/ddl/create_index.rs index 6e8f34e0..3440581f 100644 --- a/src/execution/ddl/create_index.rs +++ b/src/execution/ddl/create_index.rs @@ -137,7 +137,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateIndex { let Some(tuple_pk) = arena.result_tuple().pk.clone() else { continue; }; - with_projection_tmp_value(arena, None, &column_exprs, |arena, value| { + with_projection_tmp_value(arena, plan_arena, None, &column_exprs, |arena, value| { let mut state = arena.local_state(plan_arena); let (transaction, table_codec) = state.transaction_codec_mut(); let index = Index::new(index_id, &value, ty); diff --git a/src/execution/dml/analyze.rs b/src/execution/dml/analyze.rs index 99d33fb9..4d0b3f48 100644 --- a/src/execution/dml/analyze.rs +++ b/src/execution/dml/analyze.rs @@ -114,7 +114,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Analyze { while arena.next_tuple(input, plan_arena)? { for State { exprs, builder, .. } in builders.iter_mut() { - with_projection_tmp_value(arena, None, exprs, |_, value| builder.append(value))?; + with_projection_tmp_value(arena, plan_arena, None, exprs, |_, value| { + builder.append(value) + })?; } } let mut state = arena.local_state(plan_arena); diff --git a/src/execution/dml/copy_to_file.rs b/src/execution/dml/copy_to_file.rs index af946131..a79b40b2 100644 --- a/src/execution/dml/copy_to_file.rs +++ b/src/execution/dml/copy_to_file.rs @@ -162,7 +162,7 @@ mod tests { db.run("insert into t1 values (2, 2.0, 'fooo')")?.done()?; db.run("insert into t1 values (3, 2.1, 'Kite')")?.done()?; - let plan_arena = crate::planner::PlanArena::new(db.state.table_arena()); + let mut plan_arena = crate::planner::PlanArena::new(db.state.table_arena()); let transaction = db.storage.transaction()?; let table = transaction .table(db.state.table_cache(), "t1".to_string().into())? @@ -174,7 +174,7 @@ mod tests { "t1".to_string().into(), table, true, - &plan_arena, + &mut plan_arena, )?, column_names: Default::default(), input: None, diff --git a/src/execution/dml/delete.rs b/src/execution/dml/delete.rs index 10005b67..158002c0 100644 --- a/src/execution/dml/delete.rs +++ b/src/execution/dml/delete.rs @@ -98,7 +98,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Delete { }; for (index_id, index_ty, exprs) in index_templates.iter() { - with_projection_tmp_value(arena, None, exprs, |arena, value| { + with_projection_tmp_value(arena, plan_arena, None, exprs, |arena, value| { let mut state = arena.local_state(plan_arena); let (transaction, table_codec) = state.transaction_codec_mut(); transaction.del_index( diff --git a/src/execution/dml/insert.rs b/src/execution/dml/insert.rs index ec90c883..4e91952a 100644 --- a/src/execution/dml/insert.rs +++ b/src/execution/dml/insert.rs @@ -149,7 +149,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Insert { tuple_map.remove(&Self::column_key(column, self.is_mapping_by_name)); if value.is_none() { - value = column.default_value()?; + value = column.default_value(plan_arena)?; } value.unwrap_or(DataValue::Null) }; @@ -168,12 +168,18 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Insert { for (index_meta, exprs) in table_snapshot.index_metas.iter() { let index_meta = plan_arena.index(*index_meta); let tuple_id = tuple.pk.as_ref().ok_or(DatabaseError::PrimaryKeyNotFound)?; - with_projection_tmp_value(arena, Some(&tuple), exprs, |arena, value| { - let mut state = arena.local_state(plan_arena); - let (transaction, table_codec) = state.transaction_codec_mut(); - let index = Index::new(index_meta.id, &value, index_meta.ty); - transaction.add_index(table_codec, &self.table_name, index, tuple_id) - })?; + with_projection_tmp_value( + arena, + plan_arena, + Some(&tuple), + exprs, + |arena, value| { + let mut state = arena.local_state(plan_arena); + let (transaction, table_codec) = state.transaction_codec_mut(); + let index = Index::new(index_meta.id, &value, index_meta.ty); + transaction.add_index(table_codec, &self.table_name, index, tuple_id) + }, + )?; } let mut state = arena.local_state(plan_arena); let (transaction, table_codec) = state.transaction_codec_mut(); diff --git a/src/execution/dml/update.rs b/src/execution/dml/update.rs index df1a22b0..963b84fc 100644 --- a/src/execution/dml/update.rs +++ b/src/execution/dml/update.rs @@ -18,10 +18,9 @@ use crate::execution::{ build_read, with_projection_tmp_value, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; -use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; use crate::planner::operator::update::UpdateOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::index::{Index, IndexMeta, IndexType}; use crate::types::tuple::{Schema, Tuple}; @@ -34,7 +33,7 @@ use std::{ pub struct Update { table_name: TableName, - value_exprs: Vec<(ColumnRef, ScalarExpression)>, + value_exprs: Vec<(ColumnRef, ExprRef)>, input_schema: Schema, input_plan: LogicalPlan, input: Option, @@ -167,7 +166,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { continue; } - with_projection_tmp_value(arena, None, exprs, |_, value| { + with_projection_tmp_value(arena, plan_arena, None, exprs, |_, value| { old_index_values.push((index_offset, value)); Ok(()) })?; @@ -177,7 +176,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { continue; }; if let Some(expr) = exprs_map.get(&column_id) { - let value = expr.eval(Some(arena.result_tuple()))?; + let value = plan_arena + .expression(*expr) + .eval(plan_arena, Some(arena.result_tuple()))?; arena.result_tuple_mut().values[i] = value; } } @@ -201,7 +202,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { let index_meta = plan_arena.index(*index_meta); let index_id = index_meta.id; let index_ty = index_meta.ty; - with_projection_tmp_value(arena, None, exprs, |arena, value| { + with_projection_tmp_value(arena, plan_arena, None, exprs, |arena, value| { if !primary_key_changed && old_value == value { return Ok(()); } diff --git a/src/execution/dql/aggregate/hash_agg.rs b/src/execution/dql/aggregate/hash_agg.rs index 32622451..54bf56e1 100644 --- a/src/execution/dql/aggregate/hash_agg.rs +++ b/src/execution/dql/aggregate/hash_agg.rs @@ -19,9 +19,8 @@ use crate::execution::dql::aggregate::{ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::planner::operator::aggregate::AggregateOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; use std::collections::hash_map::IntoIter as HashMapIntoIter; @@ -30,8 +29,8 @@ use std::collections::HashMap; type HashAggOutput = HashMapIntoIter, Vec>>; pub struct HashAggExecutor { - agg_calls: Vec, - groupby_exprs: Vec, + agg_calls: Vec, + groupby_exprs: Vec, input: ExecId, output: Option, } @@ -78,14 +77,14 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashAggExecutor { let tuple = arena.result_tuple(); group_keys.clear(); for expr in &self.groupby_exprs { - group_keys.push(expr.eval(Some(tuple))?); + group_keys.push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); } if let Some(accs) = group_hash_accs.get_mut(group_keys.as_slice()) { - update_accumulators(accs, &self.agg_calls, tuple)?; + update_accumulators(accs, &self.agg_calls, tuple, plan_arena)?; } else { - let mut accs = create_accumulators(&self.agg_calls)?; - update_accumulators(&mut accs, &self.agg_calls, tuple)?; + let mut accs = create_accumulators(&self.agg_calls, plan_arena)?; + update_accumulators(&mut accs, &self.agg_calls, tuple, plan_arena)?; group_hash_accs.insert(group_keys.clone(), accs); } } @@ -174,15 +173,19 @@ mod test { }), Childrens::None, ); + let groupby_expr = + plan_arena.alloc_expression(ScalarExpression::column_expr(t1_schema[0], 0)); + let agg_arg = plan_arena.alloc_expression(ScalarExpression::column_expr(t1_schema[1], 1)); + let agg_call = plan_arena.alloc_expression(ScalarExpression::AggCall { + distinct: false, + kind: AggKind::Sum, + args: vec![agg_arg], + ty: LogicalType::Integer, + }); let plan = LogicalPlan::new( Operator::Aggregate(AggregateOperator { - groupby_exprs: vec![ScalarExpression::column_expr(t1_schema[0], 0)], - agg_calls: vec![ScalarExpression::AggCall { - distinct: false, - kind: AggKind::Sum, - args: vec![ScalarExpression::column_expr(t1_schema[1], 1)], - ty: LogicalType::Integer, - }], + groupby_exprs: vec![groupby_expr], + agg_calls: vec![agg_call], is_distinct: false, force_spill: false, }), diff --git a/src/execution/dql/aggregate/mod.rs b/src/execution/dql/aggregate/mod.rs index 803cf65a..3c27e46d 100644 --- a/src/execution/dql/aggregate/mod.rs +++ b/src/execution/dql/aggregate/mod.rs @@ -29,6 +29,7 @@ use crate::execution::dql::aggregate::sum::{DistinctSumAccumulator, SumAccumulat use crate::expression::agg::AggKind; use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use std::borrow::Cow; @@ -68,14 +69,15 @@ pub(crate) fn create_accumulator( #[inline] pub(crate) fn create_accumulators( - exprs: &[ScalarExpression], + exprs: &[ExprRef], + arena: &PlanArena<'_>, ) -> Result>, DatabaseError> { exprs .iter() .map(|expr| { let ScalarExpression::AggCall { kind, ty, distinct, .. - } = expr + } = arena.expression(*expr) else { unreachable!("create_accumulators called with non-aggregate expression {expr}") }; @@ -86,11 +88,12 @@ pub(crate) fn create_accumulators( pub(crate) fn update_accumulators( accs: &mut [Box], - agg_calls: &[ScalarExpression], + agg_calls: &[ExprRef], tuple: &Tuple, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { for (acc, expr) in accs.iter_mut().zip(agg_calls.iter()) { - let ScalarExpression::AggCall { args, .. } = expr else { + let ScalarExpression::AggCall { args, .. } = arena.expression(*expr) else { unreachable!() }; if args.len() > 1 { @@ -99,7 +102,7 @@ pub(crate) fn update_accumulators( .to_string(), )); } - let value = args[0].eval(Some(tuple))?; + let value = arena.expression(args[0]).eval(arena, Some(tuple))?; acc.update_value(&value)?; } Ok(()) diff --git a/src/execution/dql/aggregate/simple_agg.rs b/src/execution/dql/aggregate/simple_agg.rs index 74a5f15d..8b6202aa 100644 --- a/src/execution/dql/aggregate/simple_agg.rs +++ b/src/execution/dql/aggregate/simple_agg.rs @@ -19,10 +19,10 @@ use crate::execution::{ }; use crate::expression::ScalarExpression; use crate::planner::operator::aggregate::AggregateOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; pub struct SimpleAggExecutor { - agg_calls: Vec, + agg_calls: Vec, input: ExecId, returned: bool, } @@ -57,12 +57,12 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for SimpleAggExecutor { return Ok(()); } - let mut accs = create_accumulators(&self.agg_calls)?; + let mut accs = create_accumulators(&self.agg_calls, plan_arena)?; while arena.next_tuple(self.input, plan_arena)? { let tuple = arena.result_tuple(); for (acc, expr) in accs.iter_mut().zip(self.agg_calls.iter()) { - let ScalarExpression::AggCall { args, .. } = expr else { + let ScalarExpression::AggCall { args, .. } = plan_arena.expression(*expr) else { unreachable!() }; if args.len() > 1 { @@ -72,7 +72,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for SimpleAggExecutor { )); } - let value = args[0].eval(Some(tuple))?; + let value = plan_arena + .expression(args[0]) + .eval(plan_arena, Some(tuple))?; acc.update_value(&value)?; } } diff --git a/src/execution/dql/aggregate/stream_agg.rs b/src/execution/dql/aggregate/stream_agg.rs index 5f281de6..d120d705 100644 --- a/src/execution/dql/aggregate/stream_agg.rs +++ b/src/execution/dql/aggregate/stream_agg.rs @@ -19,17 +19,16 @@ use crate::execution::dql::aggregate::{ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::planner::operator::aggregate::AggregateOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; use std::mem; // The optimizer selects this executor only when equal group keys are contiguous in the input. pub struct StreamAggExecutor { - agg_calls: Vec, - groupby_exprs: Vec, + agg_calls: Vec, + groupby_exprs: Vec, group_keys: Option>, accs: Vec>, input: ExecId, @@ -87,21 +86,21 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for StreamAggExecutor { let tuple = arena.result_tuple(); let mut group_keys = Vec::with_capacity(self.groupby_exprs.len()); for expr in &self.groupby_exprs { - group_keys.push(expr.eval(Some(tuple))?); + group_keys.push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); } match &mut self.group_keys { None => { - self.accs = create_accumulators(&self.agg_calls)?; - update_accumulators(&mut self.accs, &self.agg_calls, tuple)?; + self.accs = create_accumulators(&self.agg_calls, plan_arena)?; + update_accumulators(&mut self.accs, &self.agg_calls, tuple, plan_arena)?; self.group_keys = Some(group_keys); } Some(current_keys) if current_keys == &group_keys => { - update_accumulators(&mut self.accs, &self.agg_calls, tuple)?; + update_accumulators(&mut self.accs, &self.agg_calls, tuple, plan_arena)?; } Some(current_keys) => { - let mut next_accs = create_accumulators(&self.agg_calls)?; - update_accumulators(&mut next_accs, &self.agg_calls, tuple)?; + let mut next_accs = create_accumulators(&self.agg_calls, plan_arena)?; + update_accumulators(&mut next_accs, &self.agg_calls, tuple, plan_arena)?; mem::swap(current_keys, &mut group_keys); let current_accs = mem::replace(&mut self.accs, next_accs); write_aggregate_output(arena.result_tuple_mut(), current_accs, group_keys)?; @@ -153,23 +152,23 @@ mod tests { }), Childrens::None, ); - let value = ScalarExpression::column_expr(columns[1], 1); + let group = plan_arena.alloc_expression(ScalarExpression::column_expr(columns[0], 0)); + let value = plan_arena.alloc_expression(ScalarExpression::column_expr(columns[1], 1)); + let sum = plan_arena.alloc_expression(ScalarExpression::AggCall { + distinct: false, + kind: AggKind::Sum, + args: vec![value], + ty: LogicalType::Integer, + }); + let count = plan_arena.alloc_expression(ScalarExpression::AggCall { + distinct: false, + kind: AggKind::Count, + args: vec![value], + ty: LogicalType::Integer, + }); let operator = AggregateOperator { - groupby_exprs: vec![ScalarExpression::column_expr(columns[0], 0)], - agg_calls: vec![ - ScalarExpression::AggCall { - distinct: false, - kind: AggKind::Sum, - args: vec![value.clone()], - ty: LogicalType::Integer, - }, - ScalarExpression::AggCall { - distinct: false, - kind: AggKind::Count, - args: vec![value], - ty: LogicalType::Integer, - }, - ], + groupby_exprs: vec![group], + agg_calls: vec![sum, count], is_distinct: false, force_spill: false, }; @@ -213,7 +212,9 @@ mod tests { Childrens::None, ); let operator = AggregateOperator { - groupby_exprs: vec![ScalarExpression::column_expr(column, 0)], + groupby_exprs: vec![ + plan_arena.alloc_expression(ScalarExpression::column_expr(column, 0)) + ], agg_calls: Vec::new(), is_distinct: false, force_spill: false, diff --git a/src/execution/dql/aggregate/stream_distinct.rs b/src/execution/dql/aggregate/stream_distinct.rs index 968f5988..7b1d3148 100644 --- a/src/execution/dql/aggregate/stream_distinct.rs +++ b/src/execution/dql/aggregate/stream_distinct.rs @@ -16,16 +16,15 @@ use crate::errors::DatabaseError; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; use crate::planner::operator::aggregate::AggregateOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::tuple::Tuple; use crate::types::value::DataValue; pub struct StreamDistinctExecutor { - groupby_exprs: Vec, + groupby_exprs: Vec, input: ExecId, last_keys: Option>, scratch: Tuple, @@ -67,7 +66,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for StreamDistinctExecutor { let group_keys = self .groupby_exprs .iter() - .map(|expr| expr.eval(Some(tuple))) + .map(|expr| plan_arena.expression(*expr).eval(plan_arena, Some(tuple))) .try_collect()?; if self.last_keys.as_ref() != Some(&group_keys) { @@ -161,7 +160,9 @@ mod tests { Childrens::None, ); let agg = AggregateOperator { - groupby_exprs: vec![ScalarExpression::column_expr(schema_ref[0], 0)], + groupby_exprs: vec![ + plan_arena.alloc_expression(ScalarExpression::column_expr(schema_ref[0], 0)) + ], agg_calls: vec![], is_distinct: true, force_spill: false, @@ -216,8 +217,8 @@ mod tests { ); let agg = AggregateOperator { groupby_exprs: vec![ - ScalarExpression::column_expr(schema_ref[0], 0), - ScalarExpression::column_expr(schema_ref[1], 1), + plan_arena.alloc_expression(ScalarExpression::column_expr(schema_ref[0], 0)), + plan_arena.alloc_expression(ScalarExpression::column_expr(schema_ref[1], 1)), ], agg_calls: vec![], is_distinct: true, diff --git a/src/execution/dql/external_sort.rs b/src/execution/dql/external_sort.rs index c79501d0..675060eb 100644 --- a/src/execution/dql/external_sort.rs +++ b/src/execution/dql/external_sort.rs @@ -93,7 +93,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ExternalSort { let mut runs = Vec::new(); while arena.next_tuple(self.input, plan_arena)? { let tuple = mem::take(arena.result_tuple_mut()); - if let Some(segment) = rows.push(SortRow::new(sort_fields, tuple)?)? { + if let Some(segment) = rows.push(SortRow::new(sort_fields, tuple, plan_arena)?)? { runs.push(Run::new(segment, 1)); } } @@ -299,10 +299,10 @@ mod test { ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(), )); let sort_fields = vec![SortField { - expr: ScalarExpression::ColumnRef { + expr: plan_arena.alloc_expression(ScalarExpression::ColumnRef { column: sort_column, position: 0, - }, + }), asc: true, nulls_first: false, }]; @@ -311,7 +311,11 @@ mod test { .limit(4, usize::MAX) .on_flush(|rows| sort_segment(&sort_fields, rows)); for value in [DataValue::Int32(2), DataValue::Null, DataValue::Int32(1)] { - let _ = rows.push(SortRow::new(&sort_fields, Tuple::new(None, vec![value]))?)?; + let _ = rows.push(SortRow::new( + &sort_fields, + Tuple::new(None, vec![value]), + &plan_arena, + )?)?; } let values = finish_sort(rows, Vec::new(), &sort_fields, 2)? @@ -334,10 +338,10 @@ mod test { ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(), )); let sort_fields = vec![SortField { - expr: ScalarExpression::ColumnRef { + expr: plan_arena.alloc_expression(ScalarExpression::ColumnRef { column: sort_column, position: 0, - }, + }), asc: false, nulls_first: true, }]; @@ -362,7 +366,7 @@ mod test { Some(DataValue::Int32(sequence as i32)), vec![value, DataValue::Int32(sequence as i32)], ); - if let Some(segment) = rows.push(SortRow::new(&sort_fields, tuple)?)? { + if let Some(segment) = rows.push(SortRow::new(&sort_fields, tuple, &plan_arena)?)? { runs.push(Run::new(segment, 1)); } } @@ -411,18 +415,18 @@ mod test { )); let sort_fields = vec![ SortField { - expr: ScalarExpression::ColumnRef { + expr: plan_arena.alloc_expression(ScalarExpression::ColumnRef { column: key_column, position: 0, - }, + }), asc: false, nulls_first: true, }, SortField { - expr: ScalarExpression::ColumnRef { + expr: plan_arena.alloc_expression(ScalarExpression::ColumnRef { column: sequence_column, position: 1, - }, + }), asc: true, nulls_first: false, }, @@ -439,7 +443,7 @@ mod test { }; let sequence = DataValue::Int32(sequence as i32); let tuple = Tuple::new(Some(sequence.clone()), vec![key, sequence]); - if let Some(segment) = rows.push(SortRow::new(&sort_fields, tuple)?)? { + if let Some(segment) = rows.push(SortRow::new(&sort_fields, tuple, &plan_arena)?)? { runs.push(Run::new(segment, 1)); } } diff --git a/src/execution/dql/filter.rs b/src/execution/dql/filter.rs index 6160220a..cbf80648 100644 --- a/src/execution/dql/filter.rs +++ b/src/execution/dql/filter.rs @@ -16,12 +16,11 @@ use crate::errors::DatabaseError; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::planner::operator::filter::FilterOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; pub struct Filter { - predicate: ScalarExpression, + predicate: ExprRef, input: ExecId, } @@ -52,7 +51,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Filter { return Ok(()); }; let tuple = arena.result_tuple(); - if self.predicate.eval(Some(tuple))?.is_true()? { + if plan_arena + .expression(self.predicate) + .eval(plan_arena, Some(tuple))? + .is_true()? + { arena.resume(); return Ok(()); } diff --git a/src/execution/dql/function_scan.rs b/src/execution/dql/function_scan.rs index 4a54f961..2e43d6f5 100644 --- a/src/execution/dql/function_scan.rs +++ b/src/execution/dql/function_scan.rs @@ -52,11 +52,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for FunctionScan { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - _: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut crate::planner::PlanArena<'a>, ) -> Result<(), DatabaseError> { if self.iter.is_none() { let TableFunction { args, catalog } = &self.table_function; - self.iter = Some(catalog.inner.eval(args)?); + self.iter = Some(catalog.inner.eval(args, plan_arena)?); } let tuple = self.iter.as_mut().and_then(Iterator::next).transpose()?; diff --git a/src/execution/dql/join/hash/full_join.rs b/src/execution/dql/join/hash/full_join.rs index de95a531..6916c9ce 100644 --- a/src/execution/dql/join/hash/full_join.rs +++ b/src/execution/dql/join/hash/full_join.rs @@ -18,7 +18,7 @@ use crate::execution::dql::join::hash::{ }; use crate::execution::dql::join::hash_join::BuildState; use crate::execution::dql::join::RowBitmap; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; @@ -33,7 +33,8 @@ impl JoinProbeState for FullJoinState { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { if probe_state.emitted_unmatched { @@ -68,7 +69,7 @@ impl JoinProbeState for FullJoinState { if let Some(filter_expr) = filter_expr { let full_values = SplitTupleRef::from_slices(values, &probe_state.probe_tuple.values); - if !filter(&full_values, filter_expr)? { + if !filter(&full_values, filter_expr, plan_arena)? { probe_state.has_filtered = true; self.bits.insert(*i); return Ok(Some(Self::full_right_row( @@ -100,7 +101,8 @@ impl JoinProbeState for FullJoinState { fn left_drop_next( &mut self, left_drop_state: &mut LeftDropState, - _filter_expr: Option<&ScalarExpression>, + _filter_expr: Option<&ExprRef>, + _plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { let full_schema_len = self.right_schema_len + self.left_schema_len; diff --git a/src/execution/dql/join/hash/inner_join.rs b/src/execution/dql/join/hash/inner_join.rs index 4eb62931..74ea4574 100644 --- a/src/execution/dql/join/hash/inner_join.rs +++ b/src/execution/dql/join/hash/inner_join.rs @@ -15,7 +15,7 @@ use crate::errors::DatabaseError; use crate::execution::dql::join::hash::{filter, JoinProbeState, ProbeState}; use crate::execution::dql::join::hash_join::BuildState; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::{SplitTupleRef, Tuple}; pub(crate) struct InnerJoinState; @@ -25,7 +25,8 @@ impl JoinProbeState for InnerJoinState { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { probe_state.finished = true; @@ -45,7 +46,7 @@ impl JoinProbeState for InnerJoinState { if let Some(filter_expr) = filter_expr { let full_values = SplitTupleRef::from_slices(values, &probe_state.probe_tuple.values); - if !filter(&full_values, filter_expr)? { + if !filter(&full_values, filter_expr, plan_arena)? { continue; } } diff --git a/src/execution/dql/join/hash/left_join.rs b/src/execution/dql/join/hash/left_join.rs index 7d3a123e..103d536e 100644 --- a/src/execution/dql/join/hash/left_join.rs +++ b/src/execution/dql/join/hash/left_join.rs @@ -18,7 +18,7 @@ use crate::execution::dql::join::hash::{ }; use crate::execution::dql::join::hash_join::BuildState; use crate::execution::dql::join::RowBitmap; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; @@ -33,7 +33,8 @@ impl JoinProbeState for LeftJoinState { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { probe_state.finished = true; @@ -52,7 +53,7 @@ impl JoinProbeState for LeftJoinState { if let Some(filter_expr) = filter_expr { let full_values = SplitTupleRef::from_slices(values, &probe_state.probe_tuple.values); - if !filter(&full_values, filter_expr)? { + if !filter(&full_values, filter_expr, plan_arena)? { probe_state.has_filtered = true; self.bits.insert(*i); continue; @@ -80,7 +81,8 @@ impl JoinProbeState for LeftJoinState { fn left_drop_next( &mut self, left_drop_state: &mut LeftDropState, - _filter_expr: Option<&ScalarExpression>, + _filter_expr: Option<&ExprRef>, + _plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { let full_schema_len = self.right_schema_len + self.left_schema_len; diff --git a/src/execution/dql/join/hash/mod.rs b/src/execution/dql/join/hash/mod.rs index b5d832fe..2ead9d32 100644 --- a/src/execution/dql/join/hash/mod.rs +++ b/src/execution/dql/join/hash/mod.rs @@ -24,7 +24,7 @@ use crate::execution::dql::join::hash::left_join::LeftJoinState; use crate::execution::dql::join::hash::right_join::RightJoinState; use crate::execution::dql::join::hash_join::BuildState; use crate::execution::dql::sort::BumpVec; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::{Tuple, TupleLike}; use crate::types::value::DataValue; use std::collections::hash_map::IntoIter as HashMapIntoIter; @@ -54,13 +54,15 @@ pub(crate) trait JoinProbeState { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError>; fn left_drop_next( &mut self, _left_drop_state: &mut LeftDropState, - _filter_expr: Option<&ScalarExpression>, + _filter_expr: Option<&ExprRef>, + _plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { Ok(None) } @@ -78,20 +80,21 @@ impl JoinProbeState for JoinProbeStateImpl { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { match self { JoinProbeStateImpl::Inner(state) => { - state.probe_next(probe_state, build_state, filter_expr) + state.probe_next(probe_state, build_state, filter_expr, plan_arena) } JoinProbeStateImpl::Left(state) => { - state.probe_next(probe_state, build_state, filter_expr) + state.probe_next(probe_state, build_state, filter_expr, plan_arena) } JoinProbeStateImpl::Right(state) => { - state.probe_next(probe_state, build_state, filter_expr) + state.probe_next(probe_state, build_state, filter_expr, plan_arena) } JoinProbeStateImpl::Full(state) => { - state.probe_next(probe_state, build_state, filter_expr) + state.probe_next(probe_state, build_state, filter_expr, plan_arena) } } } @@ -99,22 +102,35 @@ impl JoinProbeState for JoinProbeStateImpl { fn left_drop_next( &mut self, left_drop_state: &mut LeftDropState, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { match self { - JoinProbeStateImpl::Inner(state) => state.left_drop_next(left_drop_state, filter_expr), - JoinProbeStateImpl::Left(state) => state.left_drop_next(left_drop_state, filter_expr), - JoinProbeStateImpl::Right(state) => state.left_drop_next(left_drop_state, filter_expr), - JoinProbeStateImpl::Full(state) => state.left_drop_next(left_drop_state, filter_expr), + JoinProbeStateImpl::Inner(state) => { + state.left_drop_next(left_drop_state, filter_expr, plan_arena) + } + JoinProbeStateImpl::Left(state) => { + state.left_drop_next(left_drop_state, filter_expr, plan_arena) + } + JoinProbeStateImpl::Right(state) => { + state.left_drop_next(left_drop_state, filter_expr, plan_arena) + } + JoinProbeStateImpl::Full(state) => { + state.left_drop_next(left_drop_state, filter_expr, plan_arena) + } } } } pub(crate) fn filter( values: &T, - filter_expr: &ScalarExpression, + filter_expr: &ExprRef, + plan_arena: &PlanArena<'_>, ) -> Result { - match &filter_expr.eval(Some(values as &dyn TupleLike))? { + match &plan_arena + .expression(*filter_expr) + .eval(plan_arena, Some(values as &dyn TupleLike))? + { DataValue::Boolean(false) | DataValue::Null => Ok(false), DataValue::Boolean(true) => Ok(true), _ => Err(DatabaseError::InvalidType), diff --git a/src/execution/dql/join/hash/right_join.rs b/src/execution/dql/join/hash/right_join.rs index c226feb3..b2548c5d 100644 --- a/src/execution/dql/join/hash/right_join.rs +++ b/src/execution/dql/join/hash/right_join.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::execution::dql::join::hash::full_join::FullJoinState; use crate::execution::dql::join::hash::{filter, JoinProbeState, ProbeState}; use crate::execution::dql::join::hash_join::BuildState; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::{SplitTupleRef, Tuple}; pub(crate) struct RightJoinState { @@ -28,7 +28,8 @@ impl JoinProbeState for RightJoinState { &mut self, probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, - filter_expr: Option<&ScalarExpression>, + filter_expr: Option<&ExprRef>, + plan_arena: &PlanArena<'_>, ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { if probe_state.emitted_unmatched { @@ -63,7 +64,7 @@ impl JoinProbeState for RightJoinState { if let Some(filter_expr) = filter_expr { let full_values = SplitTupleRef::from_slices(values, &probe_state.probe_tuple.values); - if !filter(&full_values, filter_expr)? { + if !filter(&full_values, filter_expr, plan_arena)? { probe_state.has_filtered = true; continue; } diff --git a/src/execution/dql/join/hash_join.rs b/src/execution/dql/join/hash_join.rs index 99bb9d66..15e35101 100644 --- a/src/execution/dql/join/hash_join.rs +++ b/src/execution/dql/join/hash_join.rs @@ -25,9 +25,8 @@ use crate::execution::dql::sort::BumpVec; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::planner::operator::join::{JoinCondition, JoinOperator, JoinType}; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::tuple::Tuple; use crate::types::value::DataValue; @@ -38,9 +37,9 @@ use std::mem::{self, transmute}; pub struct HashJoin { state: HashJoinState, ty: JoinType, - on_left_keys: Vec, - on_right_keys: Vec, - filter: Option, + on_left_keys: Vec, + on_right_keys: Vec, + filter: Option, left_schema_len: usize, right_schema_len: usize, left_input_plan: LogicalPlan, @@ -128,13 +127,14 @@ impl HashJoin { } fn eval_keys( - on_keys: &[ScalarExpression], + on_keys: &[ExprRef], tuple: &Tuple, build_buf: &mut BumpVec<'_, DataValue>, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result<(), DatabaseError> { build_buf.clear(); for expr in on_keys { - build_buf.push(expr.eval(Some(tuple))?); + build_buf.push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); } Ok(()) } @@ -158,7 +158,7 @@ impl HashJoin { while arena.next_tuple(self.left_input, plan_arena)? { let tuple = mem::take(arena.result_tuple_mut()); - Self::eval_keys(&self.on_left_keys, &tuple, &mut build_buf)?; + Self::eval_keys(&self.on_left_keys, &tuple, &mut build_buf, plan_arena)?; match build_map.get_mut(&build_buf) { None => { @@ -284,7 +284,12 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashJoin { break true; } let tuple = mem::take(arena.result_tuple_mut()); - Self::eval_keys(&self.on_right_keys, &tuple, &mut probe_buf)?; + Self::eval_keys( + &self.on_right_keys, + &tuple, + &mut probe_buf, + plan_arena, + )?; probe_state = Some(ProbeState { is_keys_has_null: probe_buf.iter().any(DataValue::is_null), probe_tuple: tuple, @@ -305,9 +310,12 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashJoin { build_map.get_mut(&probe_buf) }; - if let Some(tuple) = - join_impl.probe_next(probe, build_state, self.filter.as_ref())? - { + if let Some(tuple) = join_impl.probe_next( + probe, + build_state, + self.filter.as_ref(), + plan_arena, + )? { if probe.finished { probe_state = None; } @@ -339,9 +347,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashJoin { mut join_impl, mut left_drop, } => { - if let Some(tuple) = - join_impl.left_drop_next(&mut left_drop, self.filter.as_ref())? - { + if let Some(tuple) = join_impl.left_drop_next( + &mut left_drop, + self.filter.as_ref(), + plan_arena, + )? { self.state = HashJoinState::LeftDrop { join_impl, left_drop, @@ -375,7 +385,7 @@ mod test { use crate::planner::operator::join::{JoinCondition, JoinOperator, JoinType}; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; - use crate::planner::{Childrens, LogicalPlan}; + use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::Storage; use crate::types::value::DataValue; @@ -399,11 +409,7 @@ mod test { fn build_join_values( arena: &mut crate::planner::PlanArena, - ) -> ( - Vec<(ScalarExpression, ScalarExpression)>, - LogicalPlan, - LogicalPlan, - ) { + ) -> (Vec<(ExprRef, ExprRef)>, LogicalPlan, LogicalPlan) { let desc = ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(); let t1_columns = vec![ @@ -418,10 +424,9 @@ mod test { arena.alloc_column(ColumnCatalog::new("c6".to_string(), true, desc.clone())), ]; - let on_keys = vec![( - ScalarExpression::column_expr(t1_columns[0], 0), - ScalarExpression::column_expr(t2_columns[0], 0), - )]; + let left_key = arena.alloc_expression(ScalarExpression::column_expr(t1_columns[0], 0)); + let right_key = arena.alloc_expression(ScalarExpression::column_expr(t2_columns[0], 0)); + let on_keys = vec![(left_key, right_key)]; let values_t1 = LogicalPlan::new( Operator::Values(ValuesOperator { @@ -682,17 +687,22 @@ mod test { let right_columns = vec![plan_arena.alloc_column(ColumnCatalog::new("rk".to_string(), true, desc.clone()))]; - let on_keys = vec![( - ScalarExpression::column_expr(left_columns[0], 0), - ScalarExpression::column_expr(right_columns[0], 0), - )]; - let filter_expr = ScalarExpression::Binary { + let left_key = + plan_arena.alloc_expression(ScalarExpression::column_expr(left_columns[0], 0)); + let right_key = + plan_arena.alloc_expression(ScalarExpression::column_expr(right_columns[0], 0)); + let on_keys = vec![(left_key, right_key)]; + let filter_left = + plan_arena.alloc_expression(ScalarExpression::column_expr(left_columns[1], 1)); + let filter_right = + plan_arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(1))); + let filter_expr = plan_arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::column_expr(left_columns[1], 1)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(1))), + left_expr: filter_left, + right_expr: filter_right, evaluator: None, ty: LogicalType::Boolean, - }; + }); let left = LogicalPlan::new( Operator::Values(ValuesOperator { diff --git a/src/execution/dql/join/nested_loop_join.rs b/src/execution/dql/join/nested_loop_join.rs index 541395e8..0847a125 100644 --- a/src/execution/dql/join/nested_loop_join.rs +++ b/src/execution/dql/join/nested_loop_join.rs @@ -22,18 +22,17 @@ use crate::execution::dql::join::RowBitmap; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; use crate::planner::operator::join::{JoinCondition, JoinOperator, JoinType}; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; /// Equivalent condition struct EqualCondition { - on_left_keys: Vec, - on_right_keys: Vec, + on_left_keys: Vec, + on_right_keys: Vec, left_len: usize, right_len: usize, } @@ -42,13 +41,22 @@ impl EqualCondition { /// Compare left tuple and right tuple on equivalent condition /// `left_tuple` must be from the [`NestedLoopJoin::left_input`] /// `right_tuple` must be from the [`NestedLoopJoin::right_input`] - fn equals(&self, left_tuple: &Tuple, right_tuple: &Tuple) -> Result { + fn equals( + &self, + left_tuple: &Tuple, + right_tuple: &Tuple, + arena: &PlanArena<'_>, + ) -> Result { if self.on_left_keys.is_empty() { return Ok(true); } for (left_expr, right_expr) in self.on_left_keys.iter().zip(self.on_right_keys.iter()) { - if left_expr.eval(Some(left_tuple))? != right_expr.eval(Some(right_tuple))? { + if arena.expression(*left_expr).eval(arena, Some(left_tuple))? + != arena + .expression(*right_expr) + .eval(arena, Some(right_tuple))? + { return Ok(false); } } @@ -70,7 +78,7 @@ pub struct NestedLoopJoin { left_input_plan: LogicalPlan, right_input_plan: LogicalPlan, ty: JoinType, - filter: Option, + filter: Option, eq_cond: EqualCondition, left_input: ExecId, state: NestedLoopJoinState, @@ -213,7 +221,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { let tuple = match ( self.filter.as_ref(), - self.eq_cond.equals(&active_left.left_tuple, &right_tuple)?, + self.eq_cond.equals( + &active_left.left_tuple, + &right_tuple, + plan_arena, + )?, ) { (None, true) if matches!(self.ty, JoinType::RightOuter) => { active_left.has_matched = true; @@ -239,7 +251,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { } else { SplitTupleRef::new(&active_left.left_tuple, &right_tuple) }; - let value = filter.eval(Some(values))?; + let value = plan_arena + .expression(*filter) + .eval(plan_arena, Some(values))?; match &value { DataValue::Boolean(true) => { let tuple = match self.ty { @@ -488,12 +502,7 @@ mod test { fn build_join_values( arena: &mut crate::planner::PlanArena, eq: bool, - ) -> ( - Vec<(ScalarExpression, ScalarExpression)>, - LogicalPlan, - LogicalPlan, - ScalarExpression, - ) { + ) -> (Vec<(ExprRef, ExprRef)>, LogicalPlan, LogicalPlan, ExprRef) { let desc = ColumnDesc::new(LogicalType::Integer, None, false, None).unwrap(); let t1_columns = vec![ @@ -510,8 +519,14 @@ mod test { let on_keys = if eq { vec![( - ScalarExpression::column_expr(t1_columns[1], 1), - ScalarExpression::column_expr(t2_columns[1], 1), + arena.alloc_expression(crate::expression::ScalarExpression::column_expr( + t1_columns[1], + 1, + )), + arena.alloc_expression(crate::expression::ScalarExpression::column_expr( + t2_columns[1], + 1, + )), )] } else { vec![] @@ -575,21 +590,27 @@ mod test { Childrens::None, ); - let filter = ScalarExpression::Binary { + let left_column = + arena.alloc_column(ColumnCatalog::new("c1".to_owned(), true, desc.clone())); + let right_column = + arena.alloc_column(ColumnCatalog::new("c4".to_owned(), true, desc.clone())); + let left_expr = arena.alloc_expression(crate::expression::ScalarExpression::column_expr( + left_column, + 0, + )); + let right_expr = arena.alloc_expression(crate::expression::ScalarExpression::column_expr( + right_column, + 3, + )); + let filter = arena.alloc_expression(crate::expression::ScalarExpression::Binary { op: crate::expression::BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new("c1".to_owned(), true, desc.clone())), - 0, - )), - right_expr: Box::new(ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new("c4".to_owned(), true, desc.clone())), - 3, - )), + left_expr, + right_expr, evaluator: Some( binary_create(Cow::Owned(LogicalType::Integer), BinaryOperator::Gt).unwrap(), ), ty: LogicalType::Boolean, - }; + }); (on_keys, values_t1, values_t2, filter) } @@ -1116,17 +1137,27 @@ mod test { let right_columns = vec![plan_arena.alloc_column(ColumnCatalog::new("rk".to_string(), true, desc.clone()))]; - let on_keys = vec![( - ScalarExpression::column_expr(left_columns[0], 0), - ScalarExpression::column_expr(right_columns[0], 0), - )]; - let filter_expr = ScalarExpression::Binary { - op: crate::expression::BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::column_expr(left_columns[1], 1)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(1))), - evaluator: None, - ty: LogicalType::Boolean, - }; + let left_key = plan_arena.alloc_expression( + crate::expression::ScalarExpression::column_expr(left_columns[0], 0), + ); + let right_key = plan_arena.alloc_expression( + crate::expression::ScalarExpression::column_expr(right_columns[0], 0), + ); + let on_keys = vec![(left_key, right_key)]; + let filter_left = plan_arena.alloc_expression( + crate::expression::ScalarExpression::column_expr(left_columns[1], 1), + ); + let filter_right = plan_arena.alloc_expression( + crate::expression::ScalarExpression::Constant(DataValue::Int32(1)), + ); + let filter_expr = + plan_arena.alloc_expression(crate::expression::ScalarExpression::Binary { + op: crate::expression::BinaryOperator::Gt, + left_expr: filter_left, + right_expr: filter_right, + evaluator: None, + ty: LogicalType::Boolean, + }); let left = LogicalPlan::new( Operator::Values(ValuesOperator { diff --git a/src/execution/dql/mark_apply.rs b/src/execution/dql/mark_apply.rs index d2fbfc38..cec3927b 100644 --- a/src/execution/dql/mark_apply.rs +++ b/src/execution/dql/mark_apply.rs @@ -141,10 +141,15 @@ impl MarkApply { fn parameterized_probe_value( &self, left_tuple: &Tuple, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result, DatabaseError> { self.op .parameterized_probe() - .map(|probe| probe.eval(Some(left_tuple))) + .map(|probe| { + plan_arena + .expression(*probe) + .eval(plan_arena, Some(left_tuple)) + }) .transpose() } @@ -158,11 +163,11 @@ impl MarkApply { MarkApplyKind::Exists => self.with_right_input( arena, plan_arena, - self.parameterized_probe_value(left_tuple)?, + self.parameterized_probe_value(left_tuple, plan_arena)?, |arena, plan_arena, right_input| { while arena.next_tuple(right_input, plan_arena)? { let right_tuple = arena.result_tuple(); - if self.exists_predicate_matched(left_tuple, right_tuple)? { + if self.exists_predicate_matched(left_tuple, right_tuple, plan_arena)? { return Ok(DataValue::Boolean(true)); } } @@ -171,7 +176,7 @@ impl MarkApply { }, ), MarkApplyKind::Quantified(MarkApplyQuantifier::Any) => { - if let Some(probe_value) = self.parameterized_probe_value(left_tuple)? { + if let Some(probe_value) = self.parameterized_probe_value(left_tuple, plan_arena)? { if !probe_value.is_null() { if self.with_right_input( arena, @@ -180,8 +185,11 @@ impl MarkApply { |arena, plan_arena, right_input| { while arena.next_tuple(right_input, plan_arena)? { let right_tuple = arena.result_tuple(); - if self.quantified_predicate_outcome(left_tuple, right_tuple)? - == QuantifiedPredicateOutcome::True + if self.quantified_predicate_outcome( + left_tuple, + right_tuple, + plan_arena, + )? == QuantifiedPredicateOutcome::True { return Ok(true); } @@ -200,8 +208,11 @@ impl MarkApply { |arena, plan_arena, right_input| { while arena.next_tuple(right_input, plan_arena)? { let right_tuple = arena.result_tuple(); - if self.quantified_predicate_outcome(left_tuple, right_tuple)? - == QuantifiedPredicateOutcome::Null + if self.quantified_predicate_outcome( + left_tuple, + right_tuple, + plan_arena, + )? == QuantifiedPredicateOutcome::Null { return Ok(true); } @@ -253,7 +264,7 @@ impl MarkApply { while arena.next_tuple(right_input, plan_arena)? { let right_tuple = arena.result_tuple(); - match self.quantified_predicate_outcome(left_tuple, right_tuple)? { + match self.quantified_predicate_outcome(left_tuple, right_tuple, plan_arena)? { QuantifiedPredicateOutcome::True => { if matches!(quantifier, MarkApplyQuantifier::Any) { return Ok(DataValue::Boolean(true)); @@ -283,11 +294,15 @@ impl MarkApply { &self, left_tuple: &Tuple, right_tuple: &Tuple, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result { let values = SplitTupleRef::new(left_tuple, right_tuple); for predicate in self.op.predicates() { - match predicate.eval(Some(values))? { + match plan_arena + .expression(*predicate) + .eval(plan_arena, Some(values))? + { DataValue::Boolean(true) => {} DataValue::Boolean(false) | DataValue::Null => return Ok(false), _ => return Err(DatabaseError::InvalidType), @@ -301,8 +316,9 @@ impl MarkApply { &self, left_tuple: &Tuple, right_tuple: &Tuple, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result { - match self.eval_predicates(left_tuple, right_tuple)? { + match self.eval_predicates(left_tuple, right_tuple, plan_arena)? { Some(DataValue::Boolean(true)) => Ok(QuantifiedPredicateOutcome::True), Some(DataValue::Boolean(false)) => Ok(QuantifiedPredicateOutcome::False), Some(DataValue::Null) => Ok(QuantifiedPredicateOutcome::Null), @@ -315,6 +331,7 @@ impl MarkApply { &self, left_tuple: &Tuple, right_tuple: &Tuple, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result, DatabaseError> { let values = SplitTupleRef::new(left_tuple, right_tuple); // probe_predicate is in predicate, always first @@ -325,14 +342,21 @@ impl MarkApply { .ok_or(DatabaseError::InvalidType)?; for predicate in correlated_predicates { - match predicate.eval(Some(values))? { + match plan_arena + .expression(*predicate) + .eval(plan_arena, Some(values))? + { DataValue::Boolean(true) => {} DataValue::Boolean(false) | DataValue::Null => return Ok(None), _ => return Err(DatabaseError::InvalidType), } } - Ok(Some(probe_predicate.eval(Some(values))?)) + Ok(Some( + plan_arena + .expression(*probe_predicate) + .eval(plan_arena, Some(values))?, + )) } } @@ -344,7 +368,7 @@ mod tests { use crate::expression::{BinaryOperator, ScalarExpression}; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; - use crate::planner::{Childrens, LogicalPlan}; + use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::{StatisticsMetaCache, Storage, TableCache, ViewCache}; use crate::types::evaluator::binary_create; @@ -413,21 +437,26 @@ mod tests { } fn build_equality_predicate( + plan_arena: &mut crate::planner::PlanArena, left_column: ColumnRef, left_position: usize, right_column: ColumnRef, right_position: usize, - ) -> Result { - Ok(ScalarExpression::Binary { + ) -> Result { + let left_expr = + plan_arena.alloc_expression(ScalarExpression::column_expr(left_column, left_position)); + let right_expr = plan_arena + .alloc_expression(ScalarExpression::column_expr(right_column, right_position)); + Ok(plan_arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Eq, - left_expr: Box::new(ScalarExpression::column_expr(left_column, left_position)), - right_expr: Box::new(ScalarExpression::column_expr(right_column, right_position)), + left_expr, + right_expr, evaluator: Some(binary_create( Cow::Owned(LogicalType::Integer), BinaryOperator::Eq, )?), ty: LogicalType::Boolean, - }) + })) } #[test] @@ -447,7 +476,7 @@ mod tests { let left_column = left.output_schema(&mut plan_arena)[0]; let right_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_column, 0, right_column, 1)?; + let predicate = build_equality_predicate(&mut plan_arena, left_column, 0, right_column, 1)?; let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -498,7 +527,7 @@ mod tests { let left_column = left.output_schema(&mut plan_arena)[0]; let right_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_column, 0, right_column, 1)?; + let predicate = build_equality_predicate(&mut plan_arena, left_column, 0, right_column, 1)?; let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -564,13 +593,16 @@ mod tests { let right_flag_column = right_schema[1]; let probe_predicate = - build_equality_predicate(left_value_column, 0, right_value_column, 2)?; - let flag_predicate = build_equality_predicate(left_flag_column, 1, right_flag_column, 3)?; + build_equality_predicate(&mut plan_arena, left_value_column, 0, right_value_column, 2)?; + let flag_predicate = + build_equality_predicate(&mut plan_arena, left_flag_column, 1, right_flag_column, 3)?; let mut op = MarkApplyOperator::new_exists( build_marker_column(&mut plan_arena), vec![probe_predicate, flag_predicate], ); - op.set_parameterized_probe(Some(ScalarExpression::column_expr(left_value_column, 0))); + let probe = + plan_arena.alloc_expression(ScalarExpression::column_expr(left_value_column, 0)); + op.set_parameterized_probe(Some(probe)); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -615,10 +647,13 @@ mod tests { ); let left_value_column = left.output_schema(&mut plan_arena)[0]; let right_value_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_value_column, 0, right_value_column, 1)?; + let predicate = + build_equality_predicate(&mut plan_arena, left_value_column, 0, right_value_column, 1)?; let mut op = MarkApplyOperator::new_in(build_marker_column(&mut plan_arena), vec![predicate]); - op.set_parameterized_probe(Some(ScalarExpression::column_expr(left_value_column, 0))); + let probe = + plan_arena.alloc_expression(ScalarExpression::column_expr(left_value_column, 0)); + op.set_parameterized_probe(Some(probe)); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -663,10 +698,13 @@ mod tests { ); let left_value_column = left.output_schema(&mut plan_arena)[0]; let right_value_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_value_column, 0, right_value_column, 1)?; + let predicate = + build_equality_predicate(&mut plan_arena, left_value_column, 0, right_value_column, 1)?; let mut op = MarkApplyOperator::new_in(build_marker_column(&mut plan_arena), vec![predicate]); - op.set_parameterized_probe(Some(ScalarExpression::column_expr(left_value_column, 0))); + op.set_parameterized_probe(Some( + plan_arena.alloc_expression(ScalarExpression::column_expr(left_value_column, 0)), + )); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -715,7 +753,7 @@ mod tests { let left_column = left.output_schema(&mut plan_arena)[0]; let right_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_column, 0, right_column, 1)?; + let predicate = build_equality_predicate(&mut plan_arena, left_column, 0, right_column, 1)?; let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -763,7 +801,7 @@ mod tests { let left_column = left.output_schema(&mut plan_arena)[0]; let right_column = right.output_schema(&mut plan_arena)[0]; - let predicate = build_equality_predicate(left_column, 0, right_column, 1)?; + let predicate = build_equality_predicate(&mut plan_arena, left_column, 0, right_column, 1)?; let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; @@ -816,17 +854,22 @@ mod tests { let right_value_column = right_schema[0]; let right_flag_column = right_schema[1]; - let probe_predicate = build_equality_predicate(left_column, 0, right_value_column, 1)?; - let correlated_predicate = ScalarExpression::Binary { + let probe_predicate = + build_equality_predicate(&mut plan_arena, left_column, 0, right_value_column, 1)?; + let correlated_left = + plan_arena.alloc_expression(ScalarExpression::column_expr(right_flag_column, 2)); + let correlated_right = + plan_arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(1))); + let correlated_predicate = plan_arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Eq, - left_expr: Box::new(ScalarExpression::column_expr(right_flag_column, 2)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(1))), + left_expr: correlated_left, + right_expr: correlated_right, evaluator: Some(binary_create( std::borrow::Cow::Owned(LogicalType::Integer), BinaryOperator::Eq, )?), ty: LogicalType::Boolean, - }; + }); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; diff --git a/src/execution/dql/projection.rs b/src/execution/dql/projection.rs index 4e3a5311..b30e7cc3 100644 --- a/src/execution/dql/projection.rs +++ b/src/execution/dql/projection.rs @@ -16,13 +16,12 @@ use crate::errors::DatabaseError; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::expression::ScalarExpression; use crate::planner::operator::project::ProjectOperator; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; pub struct Projection { - exprs: Vec, + exprs: Vec, input: ExecId, } @@ -56,7 +55,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Projection { let tuple = arena.result_tuple(); projection_tmp.reserve(self.exprs.len()); for expr in self.exprs.iter() { - projection_tmp.push(expr.eval(Some(tuple))?); + projection_tmp.push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); } std::mem::swap(&mut arena.result_tuple_mut().values, projection_tmp); Ok::<_, DatabaseError>(()) diff --git a/src/execution/dql/recursive_cte.rs b/src/execution/dql/recursive_cte.rs index 5e5751b4..a3fc5010 100644 --- a/src/execution/dql/recursive_cte.rs +++ b/src/execution/dql/recursive_cte.rs @@ -412,32 +412,34 @@ mod tests { Operator::RecursiveScan(RecursiveScanOperator { schema_ref }), Childrens::None, ); - let filter = FilterOperator::build( - ScalarExpression::Binary { - op: BinaryOperator::Lt, - left_expr: Box::new(ScalarExpression::column_expr(column, 0)), - right_expr: Box::new(DataValue::Int32(3).into()), - evaluator: Some(binary_create( - Cow::Owned(LogicalType::Integer), - BinaryOperator::Lt, - )?), - ty: LogicalType::Boolean, - }, - scan, - false, - ); + let filter_left = plan_arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let filter_right = plan_arena.alloc_expression(DataValue::Int32(3).into()); + let filter_expr = plan_arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::Lt, + left_expr: filter_left, + right_expr: filter_right, + evaluator: Some(binary_create( + Cow::Owned(LogicalType::Integer), + BinaryOperator::Lt, + )?), + ty: LogicalType::Boolean, + }); + let filter = FilterOperator::build(filter_expr, scan, false); + let project_left = plan_arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let project_right = plan_arena.alloc_expression(DataValue::Int32(1).into()); + let project_expr = plan_arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::Plus, + left_expr: project_left, + right_expr: project_right, + evaluator: Some(binary_create( + Cow::Owned(LogicalType::Integer), + BinaryOperator::Plus, + )?), + ty: LogicalType::Integer, + }); let recursive = LogicalPlan::new( Operator::Project(ProjectOperator { - exprs: vec![ScalarExpression::Binary { - op: BinaryOperator::Plus, - left_expr: Box::new(ScalarExpression::column_expr(column, 0)), - right_expr: Box::new(DataValue::Int32(1).into()), - evaluator: Some(binary_create( - Cow::Owned(LogicalType::Integer), - BinaryOperator::Plus, - )?), - ty: LogicalType::Integer, - }], + exprs: vec![project_expr], }), Childrens::Only(Box::new(filter)), ); diff --git a/src/execution/dql/sort.rs b/src/execution/dql/sort.rs index 69d24886..cb7e28a8 100644 --- a/src/execution/dql/sort.rs +++ b/src/execution/dql/sort.rs @@ -81,6 +81,7 @@ impl DerefMut for NullableVec<'_, T> { pub(crate) fn sort_tuples( sort_fields: &[SortField], tuples: &mut NullableVec<'_, (usize, Tuple)>, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result<(), DatabaseError> { // Extract the results of calculating SortFields to avoid double calculation // of data during comparison. @@ -88,7 +89,7 @@ pub(crate) fn sort_tuples( for (x, SortField { expr, .. }) in sort_fields.iter().enumerate() { for (_, tuple) in tuples.iter() { - eval_values[x].push(expr.eval(Some(tuple))?); + eval_values[x].push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); } } @@ -193,7 +194,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Sort { arena.finish(); return Ok(()); } - sort_tuples(&self.sort_fields, &mut self.rows)?; + sort_tuples(&self.sort_fields, &mut self.rows, plan_arena)?; self.rows.reverse(); } } @@ -235,8 +236,9 @@ mod test { fn sorted_rows<'a>( sort_fields: &[SortField], mut tuples: NullableVec<'a, (usize, Tuple)>, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result + 'a, DatabaseError> { - sort_tuples(sort_fields, &mut tuples)?; + sort_tuples(sort_fields, &mut tuples, plan_arena)?; let mut rows = Vec::with_capacity(tuples.len()); while let Some((_, tuple)) = tuples.pop() { rows.push(tuple); @@ -254,12 +256,13 @@ mod test { false, ColumnDesc::new(LogicalType::Integer, Some(0), false, None).unwrap(), )); + let sort_expr = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column, + position: 0, + }); let fn_sort_fields = |asc: bool, nulls_first: bool| { vec![SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column, - position: 0, - }, + expr: sort_expr, asc, nulls_first, }] @@ -351,18 +354,22 @@ mod test { fn_asc_and_nulls_first_eq(Box::new(sorted_rows( &fn_sort_fields(true, true), fn_tuples(), + &plan_arena, )?)); fn_asc_and_nulls_last_eq(Box::new(sorted_rows( &fn_sort_fields(true, false), fn_tuples(), + &plan_arena, )?)); fn_desc_and_nulls_first_eq(Box::new(sorted_rows( &fn_sort_fields(false, true), fn_tuples(), + &plan_arena, )?)); fn_desc_and_nulls_last_eq(Box::new(sorted_rows( &fn_sort_fields(false, false), fn_tuples(), + &plan_arena, )?)); Ok(()) @@ -382,22 +389,24 @@ mod test { false, ColumnDesc::new(LogicalType::Integer, Some(0), false, None).unwrap(), )); + let sort_expr_1 = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column_1, + position: 0, + }); + let sort_expr_2 = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column_2, + position: 1, + }); let fn_sort_fields = |asc_1: bool, nulls_first_1: bool, asc_2: bool, nulls_first_2: bool| { vec![ SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column_1, - position: 0, - }, + expr: sort_expr_1, asc: asc_1, nulls_first: nulls_first_1, }, SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column_2, - position: 1, - }, + expr: sort_expr_2, asc: asc_2, nulls_first: nulls_first_2, }, @@ -581,18 +590,22 @@ mod test { fn_asc_1_and_nulls_first_1_and_asc_2_and_nulls_first_2_eq(Box::new(sorted_rows( &fn_sort_fields(true, true, true, true), fn_tuples(), + &plan_arena, )?)); fn_asc_1_and_nulls_last_1_and_asc_2_and_nulls_first_2_eq(Box::new(sorted_rows( &fn_sort_fields(true, false, true, true), fn_tuples(), + &plan_arena, )?)); fn_desc_1_and_nulls_first_1_and_asc_2_and_nulls_first_2_eq(Box::new(sorted_rows( &fn_sort_fields(false, true, true, true), fn_tuples(), + &plan_arena, )?)); fn_desc_1_and_nulls_last_1_and_asc_2_and_nulls_first_2_eq(Box::new(sorted_rows( &fn_sort_fields(false, false, true, true), fn_tuples(), + &plan_arena, )?)); Ok(()) diff --git a/src/execution/dql/top_k.rs b/src/execution/dql/top_k.rs index 31fed1b0..6ceb11e0 100644 --- a/src/execution/dql/top_k.rs +++ b/src/execution/dql/top_k.rs @@ -53,6 +53,7 @@ fn top_sort<'a>( heap: &mut BTreeSet>, tuple: Tuple, keep_count: usize, + plan_arena: &crate::planner::PlanArena<'_>, ) -> Result<(), DatabaseError> { let mut full_key = BumpBytes::new_in(arena); for SortField { @@ -62,7 +63,9 @@ fn top_sort<'a>( } in sort_fields { let mut key = BumpBytes::new_in(arena); - expr.eval(Some(&tuple))? + plan_arena + .expression(*expr) + .eval(plan_arena, Some(&tuple))? .memcomparable_encode_with_null_order(&mut key, *nulls_first)?; if !asc && key.len() > 1 { for byte in key.iter_mut().skip(1) { @@ -145,6 +148,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for TopK { &mut set, mem::take(arena.result_tuple_mut()), keep_count, + plan_arena, )?; } @@ -192,12 +196,13 @@ mod test { false, ColumnDesc::new(LogicalType::Integer, Some(0), false, None).unwrap(), )); + let sort_expr = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column, + position: 0, + }); let fn_sort_fields = |asc: bool, nulls_first: bool| { vec![SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column, - position: 0, - }, + expr: sort_expr, asc, nulls_first, }] @@ -261,6 +266,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null]), 2, + &plan_arena, )?; top_sort( &arena, @@ -268,6 +274,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0)]), 2, + &plan_arena, )?; top_sort( &arena, @@ -275,6 +282,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1)]), 2, + &plan_arena, )?; fn_asc_and_nulls_first_eq(indices); @@ -286,6 +294,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null]), 2, + &plan_arena, )?; top_sort( &arena, @@ -293,6 +302,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0)]), 2, + &plan_arena, )?; top_sort( &arena, @@ -300,6 +310,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1)]), 2, + &plan_arena, )?; fn_asc_and_nulls_last_eq(indices); @@ -311,6 +322,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null]), 2, + &plan_arena, )?; top_sort( &arena, @@ -318,6 +330,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0)]), 2, + &plan_arena, )?; top_sort( &arena, @@ -325,6 +338,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1)]), 2, + &plan_arena, )?; fn_desc_and_nulls_first_eq(indices); @@ -336,6 +350,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null]), 2, + &plan_arena, )?; top_sort( &arena, @@ -343,6 +358,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0)]), 2, + &plan_arena, )?; top_sort( &arena, @@ -350,6 +366,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1)]), 2, + &plan_arena, )?; fn_desc_and_nulls_last_eq(indices); @@ -370,22 +387,24 @@ mod test { false, ColumnDesc::new(LogicalType::Integer, Some(0), false, None).unwrap(), )); + let sort_expr_1 = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column_1, + position: 0, + }); + let sort_expr_2 = plan_arena.alloc_expression(ScalarExpression::ColumnRef { + column: sort_column_2, + position: 1, + }); let fn_sort_fields = |asc_1: bool, nulls_first_1: bool, asc_2: bool, nulls_first_2: bool| { vec![ SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column_1, - position: 0, - }, + expr: sort_expr_1, asc: asc_1, nulls_first: nulls_first_1, }, SortField { - expr: ScalarExpression::ColumnRef { - column: sort_column_2, - position: 1, - }, + expr: sort_expr_2, asc: asc_2, nulls_first: nulls_first_2, }, @@ -536,6 +555,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -543,6 +563,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -550,6 +571,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -557,6 +579,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -564,6 +587,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -571,6 +595,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Int32(0)]), 4, + &plan_arena, )?; fn_asc_1_and_nulls_first_1_and_asc_2_and_nulls_first_2_eq(indices); @@ -582,6 +607,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -589,6 +615,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -596,6 +623,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -603,6 +631,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -610,6 +639,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -617,6 +647,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Int32(0)]), 4, + &plan_arena, )?; fn_asc_1_and_nulls_last_1_and_asc_2_and_nulls_first_2_eq(indices); @@ -628,6 +659,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -635,6 +667,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -642,6 +675,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -649,6 +683,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -656,6 +691,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -663,6 +699,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Int32(0)]), 4, + &plan_arena, )?; fn_desc_1_and_nulls_first_1_and_asc_2_and_nulls_first_2_eq(indices); @@ -674,6 +711,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -681,6 +719,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -688,6 +727,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Null]), 4, + &plan_arena, )?; top_sort( &arena, @@ -695,6 +735,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Null, DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -702,6 +743,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(0), DataValue::Int32(0)]), 4, + &plan_arena, )?; top_sort( &arena, @@ -709,6 +751,7 @@ mod test { &mut indices, Tuple::new(None, vec![DataValue::Int32(1), DataValue::Int32(0)]), 4, + &plan_arena, )?; fn_desc_1_and_nulls_last_1_and_asc_2_and_nulls_first_2_eq(indices); diff --git a/src/execution/dql/window.rs b/src/execution/dql/window.rs index 851da911..a4cc24ae 100644 --- a/src/execution/dql/window.rs +++ b/src/execution/dql/window.rs @@ -120,10 +120,16 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Window { } impl Window { - fn update_keys(&mut self, tuple: &Tuple) -> Result, DatabaseError> { + fn update_keys( + &mut self, + tuple: &Tuple, + plan_arena: &crate::planner::PlanArena<'_>, + ) -> Result, DatabaseError> { let mut boundary = (!self.state.started).then_some(Boundary::Partition); for (index, field) in self.sort_fields.iter().enumerate() { - let value = field.expr.eval(Some(tuple))?; + let value = plan_arena + .expression(field.expr) + .eval(plan_arena, Some(tuple))?; if self.state.started && self.state.sort_values[index] != value { if index < self.partition_by_len { boundary = Some(Boundary::Partition); @@ -146,7 +152,10 @@ impl Window { Ok(()) } - fn eval_functions(&mut self) -> Result<(), DatabaseError> { + fn eval_functions( + &mut self, + plan_arena: &crate::planner::PlanArena<'_>, + ) -> Result<(), DatabaseError> { if self.state.buffered.is_empty() { return Ok(()); } @@ -163,23 +172,29 @@ impl Window { self.state.peer_start, self.state.peer_index, output_offset + slot, + plan_arena, )?; } self.state.buffered.reverse(); Ok(()) } - fn eval(&mut self, tuple: Tuple, boundary: Option) -> Result { + fn eval( + &mut self, + tuple: Tuple, + boundary: Option, + plan_arena: &crate::planner::PlanArena<'_>, + ) -> Result { let boundary = match boundary { Some(boundary) => Some(boundary), - None => self.update_keys(&tuple)?, + None => self.update_keys(&tuple, plan_arena)?, }; if let Some(boundary) = boundary { let reached_boundary = boundary == Boundary::Partition || boundary == Boundary::Peer && self.retention == Retention::Peer; if reached_boundary && !self.state.buffered.is_empty() { self.state.pending = Some((tuple, boundary)); - self.eval_functions()?; + self.eval_functions(plan_arena)?; return Ok(true); } } @@ -197,7 +212,7 @@ impl Window { self.state.partition_rows += 1; self.state.buffered.push((row_index, tuple)); if self.retention == Retention::Row { - self.eval_functions()?; + self.eval_functions(plan_arena)?; return Ok(true); } Ok(false) @@ -228,13 +243,13 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Window { } else if arena.next_tuple(self.input, plan_arena)? { (mem::take(arena.result_tuple_mut()), None) } else { - self.eval_functions()?; + self.eval_functions(plan_arena)?; self.input_exhausted = true; output_ready = true; continue; }; - output_ready = self.eval(tuple, boundary)?; + output_ready = self.eval(tuple, boundary, plan_arena)?; } } } @@ -246,10 +261,14 @@ mod tests { use crate::catalog::ColumnRef; use crate::expression::agg::AggKind; use crate::expression::ScalarExpression; + use crate::planner::ExprRef; use crate::types::LogicalType; - fn column(position: usize) -> ScalarExpression { - ScalarExpression::column_expr(ColumnRef::new(position + 1), position) + fn column(arena: &mut crate::planner::PlanArena, position: usize) -> ExprRef { + arena.alloc_expression(ScalarExpression::column_expr( + ColumnRef::new(position + 1), + position, + )) } fn window( @@ -275,9 +294,16 @@ mod tests { #[test] fn row_materialization_streams_rows() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut plan_arena = crate::planner::PlanArena::new(&table_arena); + let partition = column(&mut plan_arena, 0); + let order = column(&mut plan_arena, 1); let mut window = window( Retention::Row, - vec![column(0).asc(), column(1).asc()], + vec![ + SortField::from(partition).asc(), + SortField::from(order).asc(), + ], 1, vec![ function::new( @@ -289,7 +315,7 @@ mod tests { ], ); for (value, expected) in [(10, [1_i64, 1]), (10, [2, 1]), (20, [3, 3])] { - window.eval(tuple(&[1, value]), None)?; + window.eval(tuple(&[1, value]), None, &plan_arena)?; let row = window.state.buffered.pop().unwrap().1; assert_eq!(row.values[2..], expected.map(DataValue::from)); } @@ -298,19 +324,26 @@ mod tests { #[test] fn peer_materialization_waits_for_peer_boundary() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut plan_arena = crate::planner::PlanArena::new(&table_arena); + let partition = column(&mut plan_arena, 0); + let value = column(&mut plan_arena, 1); let mut window = window( Retention::Peer, - vec![column(0).asc(), column(1).asc()], + vec![ + SortField::from(partition).asc(), + SortField::from(value).asc(), + ], 1, vec![function::new( WindowFunctionKind::Aggregate(AggKind::Sum), - vec![column(1)], + vec![value], LogicalType::Integer, )], ); - window.eval(tuple(&[1, 10]), None)?; - window.eval(tuple(&[1, 10]), None)?; - window.eval(tuple(&[1, 20]), None)?; + window.eval(tuple(&[1, 10]), None, &plan_arena)?; + window.eval(tuple(&[1, 10]), None, &plan_arena)?; + window.eval(tuple(&[1, 20]), None, &plan_arena)?; assert_eq!(window.state.buffered.len(), 2); assert!(window.state.pending.is_some()); Ok(()) @@ -318,19 +351,23 @@ mod tests { #[test] fn partition_materialization_waits_for_partition_boundary() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut plan_arena = crate::planner::PlanArena::new(&table_arena); + let partition = column(&mut plan_arena, 0); + let value = column(&mut plan_arena, 1); let mut window = window( Retention::Partition, - vec![column(0).asc()], + vec![SortField::from(partition).asc()], 1, vec![function::new( WindowFunctionKind::Aggregate(AggKind::Sum), - vec![column(1)], + vec![value], LogicalType::Integer, )], ); - window.eval(tuple(&[1, 3]), None)?; - window.eval(tuple(&[1, 7]), None)?; - window.eval(tuple(&[2, 5]), None)?; + window.eval(tuple(&[1, 3]), None, &plan_arena)?; + window.eval(tuple(&[1, 7]), None, &plan_arena)?; + window.eval(tuple(&[2, 5]), None, &plan_arena)?; assert_eq!(window.state.buffered.len(), 2); assert!(window.state.pending.is_some()); Ok(()) diff --git a/src/execution/dql/window/function.rs b/src/execution/dql/window/function.rs index 8484792f..e1b35550 100644 --- a/src/execution/dql/window/function.rs +++ b/src/execution/dql/window/function.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::execution::dql::aggregate::{create_accumulator, Accumulator}; use crate::expression::agg::AggKind; use crate::expression::window::WindowFunctionKind; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -34,6 +34,7 @@ pub(super) trait WindowFunction { peer_start: usize, peer_index: usize, output_position: usize, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError>; } @@ -47,6 +48,7 @@ impl WindowFunction for RowNumber { _peer_start: usize, _peer_index: usize, output_position: usize, + _arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { for (row_index, row) in &mut rows[peer] { row.values[output_position] = DataValue::Int64((*row_index + 1) as i64); @@ -67,6 +69,7 @@ impl WindowFunction for Rank { peer_start: usize, peer_index: usize, output_position: usize, + _arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { let rank = if self.dense { peer_index + 1 @@ -83,7 +86,7 @@ impl WindowFunction for Rank { struct Aggregate { kind: AggKind, ty: LogicalType, - arg: ScalarExpression, + arg: ExprRef, accumulator: Option>, } @@ -100,12 +103,13 @@ impl WindowFunction for Aggregate { _peer_start: usize, _peer_index: usize, output_position: usize, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { let Some(accumulator) = self.accumulator.as_mut() else { unreachable!() }; for (_, row) in &rows[peer.clone()] { - accumulator.update_value(&self.arg.eval(Some(row))?)?; + accumulator.update_value(&arena.expression(self.arg).eval(arena, Some(row))?)?; } accumulator.evaluate()?; let result = accumulator.result(); @@ -118,7 +122,7 @@ impl WindowFunction for Aggregate { pub(super) fn new( kind: WindowFunctionKind, - args: Vec, + args: Vec, ty: LogicalType, ) -> Box { match kind { diff --git a/src/execution/mod.rs b/src/execution/mod.rs index 71ae7bff..b3cc8e2e 100644 --- a/src/execution/mod.rs +++ b/src/execution/mod.rs @@ -441,17 +441,18 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { } pub(crate) fn with_projection_tmp_value<'a, T: Transaction + 'a>( - arena: &mut ExecArena<'a, T>, + exec_arena: &mut ExecArena<'a, T>, + plan_arena: &crate::planner::PlanArena<'_>, tuple: Option<&dyn TupleLike>, exprs: &[ScalarExpression], f: impl FnOnce(&mut ExecArena<'a, T>, DataValue) -> Result<(), DatabaseError>, ) -> Result<(), DatabaseError> { - arena.with_projection_tmp(|arena, projection_tmp| { + exec_arena.with_projection_tmp(|exec_arena, projection_tmp| { { - let tuple = tuple.unwrap_or_else(|| arena.result_tuple() as &dyn TupleLike); + let tuple = tuple.unwrap_or_else(|| exec_arena.result_tuple() as &dyn TupleLike); projection_tmp.reserve(exprs.len()); for expr in exprs.iter() { - projection_tmp.push(expr.eval(Some(tuple))?); + projection_tmp.push(expr.eval(plan_arena, Some(tuple))?); } } @@ -459,11 +460,11 @@ pub(crate) fn with_projection_tmp_value<'a, T: Transaction + 'a>( 0 => {} 1 => { let value = projection_tmp.pop().expect("projection has one value"); - f(arena, value)?; + f(exec_arena, value)?; } _ => { let value = DataValue::Tuple(std::mem::take(projection_tmp), false); - f(arena, value)?; + f(exec_arena, value)?; } } Ok(()) diff --git a/src/execution/spill/codec.rs b/src/execution/spill/codec.rs index 233029d6..4e000688 100644 --- a/src/execution/spill/codec.rs +++ b/src/execution/spill/codec.rs @@ -15,6 +15,7 @@ use super::SpillCodec; use crate::errors::DatabaseError; use crate::planner::operator::sort::SortField; +use crate::planner::PlanArena; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use std::io::{Read, Write}; @@ -26,10 +27,14 @@ pub(crate) struct SortRow { } impl SortRow { - pub(crate) fn new(sort_fields: &[SortField], tuple: Tuple) -> Result { + pub(crate) fn new( + sort_fields: &[SortField], + tuple: Tuple, + arena: &PlanArena<'_>, + ) -> Result { let sort_values = sort_fields .iter() - .map(|field| field.expr.eval(Some(&tuple))) + .map(|field| arena.expression(field.expr).eval(arena, Some(&tuple))) .collect::>()?; Ok(Self { sort_values, tuple }) } diff --git a/src/expression/eq_col.rs b/src/expression/eq_col.rs new file mode 100644 index 00000000..0594ac8b --- /dev/null +++ b/src/expression/eq_col.rs @@ -0,0 +1,990 @@ +// Copyright 2024 KipData/KiteSQL +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use crate::catalog::ColumnRef; +use crate::errors::DatabaseError; +use crate::expression::agg::AggKind; +use crate::expression::function::scala::ScalarFunction; +use crate::expression::function::table::TableFunction; +use crate::expression::visitor::{walk_expr, ExprVisitor}; +use crate::expression::window::WindowCall; +use crate::expression::{BinaryOperator, ScalarExpression, TrimWhereField, UnaryOperator}; +use crate::planner::{ExprRef, PlanArena}; +use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef}; +use crate::types::value::DataValue; +use crate::types::LogicalType; + +pub(super) fn eq_ignore_colref_pos(lhs: ExprRef, rhs: ExprRef, arena: &PlanArena<'_>) -> bool { + EqIgnoreColRefPosVisitor::equals(lhs, rhs, arena) +} + +struct EqIgnoreColRefPosVisitor<'a, 'arena> { + rhs: ExprRef, + arena: &'a PlanArena<'arena>, + equal: bool, +} + +impl<'a, 'arena> EqIgnoreColRefPosVisitor<'a, 'arena> { + fn equals(lhs: ExprRef, rhs: ExprRef, arena: &'a PlanArena<'arena>) -> bool { + let mut visitor = Self { + rhs, + arena, + equal: true, + }; + visitor.visit(lhs, arena).is_ok() && visitor.equal + } + + fn rhs(&self) -> &ScalarExpression { + self.arena.expression(self.rhs.unpack_alias(self.arena)) + } + + fn refs_equal(&self, lhs: &[ExprRef], rhs: &[ExprRef]) -> bool { + lhs.len() == rhs.len() + && lhs + .iter() + .zip(rhs) + .all(|(lhs, rhs)| Self::equals(*lhs, *rhs, self.arena)) + } + + fn optional_refs_equal(&self, lhs: Option, rhs: Option) -> bool { + match (lhs, rhs) { + (Some(lhs), Some(rhs)) => Self::equals(lhs, rhs, self.arena), + (None, None) => true, + _ => false, + } + } +} + +impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { + fn visit(&mut self, lhs: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + let lhs = lhs.unpack_alias(arena); + self.rhs = self.rhs.unpack_alias(arena); + if lhs == self.rhs { + return Ok(()); + } + walk_expr(self, lhs, arena) + } + + fn visit_constant(&mut self, lhs: &DataValue) -> Result<(), DatabaseError> { + self.equal = matches!(self.rhs(), ScalarExpression::Constant(rhs) if lhs == rhs); + Ok(()) + } + + fn visit_column_ref(&mut self, lhs: &ColumnRef) -> Result<(), DatabaseError> { + self.equal = matches!( + self.rhs(), + ScalarExpression::ColumnRef { column: rhs, .. } + if self.arena.same_column(*lhs, *rhs) + ); + Ok(()) + } + + fn visit_type_cast( + &mut self, + lhs_expr: ExprRef, + lhs_ty: &LogicalType, + lhs_evaluator: Option<&CastEvaluatorRef>, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::TypeCast { + expr: rhs_expr, + ty: rhs_ty, + evaluator: rhs_evaluator, + } => { + lhs_ty == rhs_ty + && lhs_evaluator == rhs_evaluator.as_ref() + && Self::equals(lhs_expr, *rhs_expr, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_is_null( + &mut self, + lhs_negated: bool, + lhs_expr: ExprRef, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::IsNull { + negated: rhs_negated, + expr: rhs_expr, + } => lhs_negated == *rhs_negated && Self::equals(lhs_expr, *rhs_expr, self.arena), + _ => false, + }; + Ok(()) + } + + fn visit_unary( + &mut self, + lhs_op: &UnaryOperator, + lhs_expr: ExprRef, + lhs_evaluator: Option<&UnaryEvaluatorRef>, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Unary { + op: rhs_op, + expr: rhs_expr, + evaluator: rhs_evaluator, + ty: rhs_ty, + } => { + lhs_op == rhs_op + && lhs_evaluator == rhs_evaluator.as_ref() + && lhs_ty == rhs_ty + && Self::equals(lhs_expr, *rhs_expr, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_binary( + &mut self, + lhs_op: &BinaryOperator, + lhs_left: ExprRef, + lhs_right: ExprRef, + lhs_evaluator: Option<&BinaryEvaluatorRef>, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Binary { + op: rhs_op, + left_expr: rhs_left, + right_expr: rhs_right, + evaluator: rhs_evaluator, + ty: rhs_ty, + } => { + lhs_op == rhs_op + && lhs_evaluator == rhs_evaluator.as_ref() + && lhs_ty == rhs_ty + && Self::equals(lhs_left, *rhs_left, self.arena) + && Self::equals(lhs_right, *rhs_right, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_agg( + &mut self, + lhs_distinct: bool, + lhs_kind: &AggKind, + lhs_args: &[ExprRef], + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::AggCall { + distinct: rhs_distinct, + kind: rhs_kind, + args: rhs_args, + ty: rhs_ty, + } => { + lhs_distinct == *rhs_distinct + && lhs_kind == rhs_kind + && lhs_ty == rhs_ty + && self.refs_equal(lhs_args, rhs_args) + } + _ => false, + }; + Ok(()) + } + + fn visit_window( + &mut self, + lhs: &WindowCall, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::WindowCall(rhs) => { + lhs.function.kind == rhs.function.kind + && lhs.function.ty == rhs.function.ty + && self.refs_equal(&lhs.function.args, &rhs.function.args) + && self.refs_equal(&lhs.spec.partition_by, &rhs.spec.partition_by) + && lhs.spec.order_by.len() == rhs.spec.order_by.len() + && lhs.spec.order_by.iter().zip(&rhs.spec.order_by).all( + |(lhs_field, rhs_field)| { + lhs_field.asc == rhs_field.asc + && lhs_field.nulls_first == rhs_field.nulls_first + && Self::equals(lhs_field.expr, rhs_field.expr, self.arena) + }, + ) + } + _ => false, + }; + Ok(()) + } + + fn visit_in( + &mut self, + lhs_negated: bool, + lhs_expr: ExprRef, + lhs_args: &[ExprRef], + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::In { + negated: rhs_negated, + expr: rhs_expr, + args: rhs_args, + } => { + lhs_negated == *rhs_negated + && Self::equals(lhs_expr, *rhs_expr, self.arena) + && self.refs_equal(lhs_args, rhs_args) + } + _ => false, + }; + Ok(()) + } + + fn visit_between( + &mut self, + lhs_negated: bool, + lhs_expr: ExprRef, + lhs_left: ExprRef, + lhs_right: ExprRef, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Between { + negated: rhs_negated, + expr: rhs_expr, + left_expr: rhs_left, + right_expr: rhs_right, + } => { + lhs_negated == *rhs_negated + && Self::equals(lhs_expr, *rhs_expr, self.arena) + && Self::equals(lhs_left, *rhs_left, self.arena) + && Self::equals(lhs_right, *rhs_right, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_substring( + &mut self, + lhs_expr: ExprRef, + lhs_for: Option, + lhs_from: Option, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::SubString { + expr: rhs_expr, + for_expr: rhs_for, + from_expr: rhs_from, + } => { + Self::equals(lhs_expr, *rhs_expr, self.arena) + && self.optional_refs_equal(lhs_for, *rhs_for) + && self.optional_refs_equal(lhs_from, *rhs_from) + } + _ => false, + }; + Ok(()) + } + + fn visit_position( + &mut self, + lhs_expr: ExprRef, + lhs_in: ExprRef, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Position { + expr: rhs_expr, + in_expr: rhs_in, + } => { + Self::equals(lhs_expr, *rhs_expr, self.arena) + && Self::equals(lhs_in, *rhs_in, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_trim( + &mut self, + lhs_expr: ExprRef, + lhs_what: Option, + lhs_where: Option<&TrimWhereField>, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Trim { + expr: rhs_expr, + trim_what_expr: rhs_what, + trim_where: rhs_where, + } => { + lhs_where == rhs_where.as_ref() + && Self::equals(lhs_expr, *rhs_expr, self.arena) + && self.optional_refs_equal(lhs_what, *rhs_what) + } + _ => false, + }; + Ok(()) + } + + fn visit_empty(&mut self) -> Result<(), DatabaseError> { + self.equal = matches!(self.rhs(), ScalarExpression::Empty); + Ok(()) + } + + fn visit_tuple( + &mut self, + lhs: &[ExprRef], + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = + matches!(self.rhs(), ScalarExpression::Tuple(rhs) if self.refs_equal(lhs, rhs)); + Ok(()) + } + + fn visit_scala_function( + &mut self, + lhs: &ScalarFunction, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = matches!( + self.rhs(), + ScalarExpression::ScalaFunction(rhs) + if lhs.summary() == rhs.summary() && self.refs_equal(&lhs.args, &rhs.args) + ); + Ok(()) + } + + fn visit_table_function( + &mut self, + lhs: &TableFunction, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = matches!( + self.rhs(), + ScalarExpression::TableFunction(rhs) + if lhs.summary() == rhs.summary() && self.refs_equal(&lhs.args, &rhs.args) + ); + Ok(()) + } + + fn visit_if( + &mut self, + lhs_condition: ExprRef, + lhs_left: ExprRef, + lhs_right: ExprRef, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::If { + condition: rhs_condition, + left_expr: rhs_left, + right_expr: rhs_right, + ty: rhs_ty, + } => { + lhs_ty == rhs_ty + && Self::equals(lhs_condition, *rhs_condition, self.arena) + && Self::equals(lhs_left, *rhs_left, self.arena) + && Self::equals(lhs_right, *rhs_right, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_if_null( + &mut self, + lhs_left: ExprRef, + lhs_right: ExprRef, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::IfNull { + left_expr: rhs_left, + right_expr: rhs_right, + ty: rhs_ty, + } => { + lhs_ty == rhs_ty + && Self::equals(lhs_left, *rhs_left, self.arena) + && Self::equals(lhs_right, *rhs_right, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_null_if( + &mut self, + lhs_left: ExprRef, + lhs_right: ExprRef, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::NullIf { + left_expr: rhs_left, + right_expr: rhs_right, + ty: rhs_ty, + } => { + lhs_ty == rhs_ty + && Self::equals(lhs_left, *rhs_left, self.arena) + && Self::equals(lhs_right, *rhs_right, self.arena) + } + _ => false, + }; + Ok(()) + } + + fn visit_coalesce( + &mut self, + lhs_exprs: &[ExprRef], + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::Coalesce { + exprs: rhs_exprs, + ty: rhs_ty, + } => lhs_ty == rhs_ty && self.refs_equal(lhs_exprs, rhs_exprs), + _ => false, + }; + Ok(()) + } + + fn visit_case_when( + &mut self, + lhs_operand: Option, + lhs_pairs: &[(ExprRef, ExprRef)], + lhs_else: Option, + lhs_ty: &LogicalType, + _arena: &PlanArena<'_>, + ) -> Result<(), DatabaseError> { + self.equal = match self.rhs() { + ScalarExpression::CaseWhen { + operand_expr: rhs_operand, + expr_pairs: rhs_pairs, + else_expr: rhs_else, + ty: rhs_ty, + } => { + lhs_ty == rhs_ty + && self.optional_refs_equal(lhs_operand, *rhs_operand) + && lhs_pairs.len() == rhs_pairs.len() + && lhs_pairs.iter().zip(rhs_pairs).all( + |((lhs_when, lhs_then), (rhs_when, rhs_then))| { + Self::equals(*lhs_when, *rhs_when, self.arena) + && Self::equals(*lhs_then, *rhs_then, self.arena) + }, + ) + && self.optional_refs_equal(lhs_else, *rhs_else) + } + _ => false, + }; + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::catalog::{ColumnCatalog, ColumnDesc}; + use crate::expression::function::scala::ArcScalarFunctionImpl; + use crate::expression::function::table::{ArcTableFunctionImpl, TableFunctionCatalog}; + use crate::expression::window::{WindowFunction, WindowFunctionKind, WindowSpec}; + use crate::expression::AliasType; + use crate::function::current_date::CurrentDate; + use crate::function::numbers::Numbers; + use crate::planner::operator::sort::SortField; + use crate::planner::TableArenaCell; + + fn assert_case( + arena: &mut PlanArena<'_>, + lhs: ScalarExpression, + rhs: ScalarExpression, + different: ScalarExpression, + ) { + let lhs = arena.alloc_expression(lhs); + let rhs = arena.alloc_expression(rhs); + let different = arena.alloc_expression(different); + + assert!( + eq_ignore_colref_pos(lhs, rhs, arena), + "lhs={lhs:?}, rhs={rhs:?}" + ); + assert!( + eq_ignore_colref_pos(rhs, lhs, arena), + "rhs={rhs:?}, lhs={lhs:?}" + ); + assert!( + eq_ignore_colref_pos(lhs, lhs, arena), + "self comparison failed: {lhs:?}" + ); + assert!( + !eq_ignore_colref_pos(lhs, different, arena), + "unexpected equality: lhs={lhs:?}, different={different:?}" + ); + assert!( + !eq_ignore_colref_pos(different, lhs, arena), + "unexpected reverse equality: different={different:?}, lhs={lhs:?}" + ); + } + + #[test] + fn compares_every_scalar_expression_variant() -> Result<(), DatabaseError> { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + + let lhs_column = arena.alloc_column(ColumnCatalog::new( + "c1".to_string(), + false, + ColumnDesc::new(LogicalType::Integer, None, false, None)?, + )); + let rhs_column = arena.alloc_column(ColumnCatalog::new( + "c1".to_string(), + true, + ColumnDesc::new(LogicalType::Bigint, None, false, None)?, + )); + let different_column = arena.alloc_column(ColumnCatalog::new( + "c2".to_string(), + false, + ColumnDesc::new(LogicalType::Integer, None, false, None)?, + )); + + let lhs_child = arena.alloc_expression(ScalarExpression::column_expr(lhs_column, 0)); + let rhs_child = arena.alloc_expression(ScalarExpression::column_expr(rhs_column, 99)); + let different_child = arena.alloc_expression(ScalarExpression::Constant(2.into())); + let lhs_one = arena.alloc_expression(ScalarExpression::Constant(1.into())); + let rhs_one = arena.alloc_expression(ScalarExpression::Constant(1.into())); + let lhs_three = arena.alloc_expression(ScalarExpression::Constant(3.into())); + let rhs_three = arena.alloc_expression(ScalarExpression::Constant(3.into())); + + assert_case( + &mut arena, + ScalarExpression::Constant(1.into()), + ScalarExpression::Constant(1.into()), + ScalarExpression::Constant(2.into()), + ); + assert_case( + &mut arena, + ScalarExpression::column_expr(lhs_column, 0), + ScalarExpression::column_expr(rhs_column, 42), + ScalarExpression::column_expr(different_column, 0), + ); + assert_case( + &mut arena, + ScalarExpression::Alias { + expr: lhs_child, + alias: AliasType::Name("lhs".to_string()), + }, + ScalarExpression::Alias { + expr: rhs_child, + alias: AliasType::Name("rhs".to_string()), + }, + ScalarExpression::Alias { + expr: different_child, + alias: AliasType::Name("lhs".to_string()), + }, + ); + assert_case( + &mut arena, + ScalarExpression::Alias { + expr: different_child, + alias: AliasType::Expr(lhs_child), + }, + ScalarExpression::Alias { + expr: different_child, + alias: AliasType::Expr(rhs_child), + }, + ScalarExpression::Alias { + expr: lhs_child, + alias: AliasType::Expr(different_child), + }, + ); + assert_case( + &mut arena, + ScalarExpression::TypeCast { + expr: lhs_child, + ty: LogicalType::Integer, + evaluator: None, + }, + ScalarExpression::TypeCast { + expr: rhs_child, + ty: LogicalType::Integer, + evaluator: None, + }, + ScalarExpression::TypeCast { + expr: rhs_child, + ty: LogicalType::Bigint, + evaluator: None, + }, + ); + assert_case( + &mut arena, + ScalarExpression::IsNull { + negated: false, + expr: lhs_child, + }, + ScalarExpression::IsNull { + negated: false, + expr: rhs_child, + }, + ScalarExpression::IsNull { + negated: true, + expr: rhs_child, + }, + ); + assert_case( + &mut arena, + ScalarExpression::Unary { + op: UnaryOperator::Minus, + expr: lhs_child, + evaluator: None, + ty: LogicalType::Integer, + }, + ScalarExpression::Unary { + op: UnaryOperator::Minus, + expr: rhs_child, + evaluator: None, + ty: LogicalType::Integer, + }, + ScalarExpression::Unary { + op: UnaryOperator::Plus, + expr: rhs_child, + evaluator: None, + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::Binary { + op: BinaryOperator::Plus, + left_expr: lhs_child, + right_expr: lhs_one, + evaluator: None, + ty: LogicalType::Integer, + }, + ScalarExpression::Binary { + op: BinaryOperator::Plus, + left_expr: rhs_child, + right_expr: rhs_one, + evaluator: None, + ty: LogicalType::Integer, + }, + ScalarExpression::Binary { + op: BinaryOperator::Minus, + left_expr: rhs_child, + right_expr: rhs_one, + evaluator: None, + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::AggCall { + distinct: true, + kind: AggKind::Sum, + args: vec![lhs_child], + ty: LogicalType::Integer, + }, + ScalarExpression::AggCall { + distinct: true, + kind: AggKind::Sum, + args: vec![rhs_child], + ty: LogicalType::Integer, + }, + ScalarExpression::AggCall { + distinct: false, + kind: AggKind::Sum, + args: vec![rhs_child], + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::In { + negated: false, + expr: lhs_child, + args: vec![lhs_one], + }, + ScalarExpression::In { + negated: false, + expr: rhs_child, + args: vec![rhs_one], + }, + ScalarExpression::In { + negated: true, + expr: rhs_child, + args: vec![rhs_one], + }, + ); + assert_case( + &mut arena, + ScalarExpression::Between { + negated: false, + expr: lhs_child, + left_expr: lhs_one, + right_expr: lhs_three, + }, + ScalarExpression::Between { + negated: false, + expr: rhs_child, + left_expr: rhs_one, + right_expr: rhs_three, + }, + ScalarExpression::Between { + negated: true, + expr: rhs_child, + left_expr: rhs_one, + right_expr: rhs_three, + }, + ); + assert_case( + &mut arena, + ScalarExpression::SubString { + expr: lhs_child, + for_expr: Some(lhs_one), + from_expr: None, + }, + ScalarExpression::SubString { + expr: rhs_child, + for_expr: Some(rhs_one), + from_expr: None, + }, + ScalarExpression::SubString { + expr: rhs_child, + for_expr: Some(rhs_one), + from_expr: Some(rhs_three), + }, + ); + assert_case( + &mut arena, + ScalarExpression::Position { + expr: lhs_child, + in_expr: lhs_one, + }, + ScalarExpression::Position { + expr: rhs_child, + in_expr: rhs_one, + }, + ScalarExpression::Position { + expr: rhs_child, + in_expr: different_child, + }, + ); + assert_case( + &mut arena, + ScalarExpression::Trim { + expr: lhs_child, + trim_what_expr: Some(lhs_one), + trim_where: Some(TrimWhereField::Both), + }, + ScalarExpression::Trim { + expr: rhs_child, + trim_what_expr: Some(rhs_one), + trim_where: Some(TrimWhereField::Both), + }, + ScalarExpression::Trim { + expr: rhs_child, + trim_what_expr: Some(rhs_one), + trim_where: Some(TrimWhereField::Leading), + }, + ); + assert_case( + &mut arena, + ScalarExpression::Empty, + ScalarExpression::Empty, + ScalarExpression::Constant(1.into()), + ); + assert_case( + &mut arena, + ScalarExpression::Tuple(vec![lhs_child, lhs_one]), + ScalarExpression::Tuple(vec![rhs_child, rhs_one]), + ScalarExpression::Tuple(vec![rhs_child]), + ); + assert_case( + &mut arena, + ScalarExpression::ScalaFunction(ScalarFunction { + args: vec![lhs_child], + inner: ArcScalarFunctionImpl(CurrentDate::new()), + }), + ScalarExpression::ScalaFunction(ScalarFunction { + args: vec![rhs_child], + inner: ArcScalarFunctionImpl(CurrentDate::new()), + }), + ScalarExpression::ScalaFunction(ScalarFunction { + args: vec![different_child], + inner: ArcScalarFunctionImpl(CurrentDate::new()), + }), + ); + assert_case( + &mut arena, + ScalarExpression::TableFunction(TableFunction { + args: vec![lhs_child], + catalog: TableFunctionCatalog { + schema: vec![], + inner: ArcTableFunctionImpl(Numbers::new()), + }, + }), + ScalarExpression::TableFunction(TableFunction { + args: vec![rhs_child], + catalog: TableFunctionCatalog { + schema: vec![], + inner: ArcTableFunctionImpl(Numbers::new()), + }, + }), + ScalarExpression::TableFunction(TableFunction { + args: vec![different_child], + catalog: TableFunctionCatalog { + schema: vec![], + inner: ArcTableFunctionImpl(Numbers::new()), + }, + }), + ); + assert_case( + &mut arena, + ScalarExpression::If { + condition: lhs_child, + left_expr: lhs_one, + right_expr: lhs_three, + ty: LogicalType::Integer, + }, + ScalarExpression::If { + condition: rhs_child, + left_expr: rhs_one, + right_expr: rhs_three, + ty: LogicalType::Integer, + }, + ScalarExpression::If { + condition: rhs_child, + left_expr: rhs_one, + right_expr: different_child, + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::IfNull { + left_expr: lhs_child, + right_expr: lhs_one, + ty: LogicalType::Integer, + }, + ScalarExpression::IfNull { + left_expr: rhs_child, + right_expr: rhs_one, + ty: LogicalType::Integer, + }, + ScalarExpression::IfNull { + left_expr: rhs_child, + right_expr: different_child, + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::NullIf { + left_expr: lhs_child, + right_expr: lhs_one, + ty: LogicalType::Integer, + }, + ScalarExpression::NullIf { + left_expr: rhs_child, + right_expr: rhs_one, + ty: LogicalType::Integer, + }, + ScalarExpression::NullIf { + left_expr: rhs_child, + right_expr: rhs_one, + ty: LogicalType::Bigint, + }, + ); + assert_case( + &mut arena, + ScalarExpression::Coalesce { + exprs: vec![lhs_child, lhs_one], + ty: LogicalType::Integer, + }, + ScalarExpression::Coalesce { + exprs: vec![rhs_child, rhs_one], + ty: LogicalType::Integer, + }, + ScalarExpression::Coalesce { + exprs: vec![rhs_child], + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::CaseWhen { + operand_expr: Some(lhs_child), + expr_pairs: vec![(lhs_child, lhs_one)], + else_expr: Some(lhs_three), + ty: LogicalType::Integer, + }, + ScalarExpression::CaseWhen { + operand_expr: Some(rhs_child), + expr_pairs: vec![(rhs_child, rhs_one)], + else_expr: Some(rhs_three), + ty: LogicalType::Integer, + }, + ScalarExpression::CaseWhen { + operand_expr: None, + expr_pairs: vec![(rhs_child, rhs_one)], + else_expr: Some(rhs_three), + ty: LogicalType::Integer, + }, + ); + assert_case( + &mut arena, + ScalarExpression::WindowCall(WindowCall { + function: WindowFunction { + kind: WindowFunctionKind::RowNumber, + args: vec![lhs_child], + ty: LogicalType::Bigint, + }, + spec: WindowSpec { + partition_by: vec![lhs_child], + order_by: vec![SortField::new(lhs_one, true, false)], + }, + }), + ScalarExpression::WindowCall(WindowCall { + function: WindowFunction { + kind: WindowFunctionKind::RowNumber, + args: vec![rhs_child], + ty: LogicalType::Bigint, + }, + spec: WindowSpec { + partition_by: vec![rhs_child], + order_by: vec![SortField::new(rhs_one, true, false)], + }, + }), + ScalarExpression::WindowCall(WindowCall { + function: WindowFunction { + kind: WindowFunctionKind::RowNumber, + args: vec![rhs_child], + ty: LogicalType::Bigint, + }, + spec: WindowSpec { + partition_by: vec![rhs_child], + order_by: vec![SortField::new(rhs_one, false, false)], + }, + }), + ); + + Ok(()) + } +} diff --git a/src/expression/evaluator.rs b/src/expression/evaluator.rs index 6bec4bf4..391c4643 100644 --- a/src/expression/evaluator.rs +++ b/src/expression/evaluator.rs @@ -15,6 +15,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::ScalarFunction; use crate::expression::{AliasType, BinaryOperator, ScalarExpression, TrimWhereField}; +use crate::planner::{ExprRef, PlanArena}; use crate::types::evaluator::binary_create; use crate::types::tuple::TupleLike; use crate::types::value::{DataValue, Utf8Type}; @@ -24,8 +25,13 @@ use std::cmp; use std::cmp::Ordering; macro_rules! eval_to_num { - ($num_expr:expr, $tuple:expr) => { - if let Some(num_i32) = $num_expr.eval($tuple)?.cast(&LogicalType::Integer)?.i32() { + ($num_expr:expr, $arena:expr, $tuple:expr) => { + if let Some(num_i32) = $arena + .expression(*$num_expr) + .eval($arena, $tuple)? + .cast(&LogicalType::Integer)? + .i32() + { num_i32 } else { return Ok(DataValue::Null); @@ -34,7 +40,11 @@ macro_rules! eval_to_num { } impl ScalarExpression { - pub fn eval(&self, tuple: Option) -> Result { + pub fn eval( + &self, + arena: &PlanArena<'_>, + tuple: Option, + ) -> Result { match self { ScalarExpression::Constant(val) => Ok(val.clone()), ScalarExpression::ColumnRef { position, .. } => { @@ -48,15 +58,15 @@ impl ScalarExpression { return Ok(DataValue::Null); }; if let AliasType::Expr(inner_expr) = alias { - inner_expr.eval(Some(tuple)) + arena.expression(*inner_expr).eval(arena, Some(tuple)) } else { - expr.eval(Some(tuple)) + arena.expression(*expr).eval(arena, Some(tuple)) } } ScalarExpression::TypeCast { expr, evaluator, .. } => { - let value = expr.eval(tuple)?; + let value = arena.expression(*expr).eval(arena, tuple)?; if let Some(evaluator) = evaluator { evaluator.eval(&value) } else { @@ -69,8 +79,8 @@ impl ScalarExpression { evaluator, .. } => { - let left = left_expr.eval(tuple)?; - let right = right_expr.eval(tuple)?; + let left = arena.expression(*left_expr).eval(arena, tuple)?; + let right = arena.expression(*right_expr).eval(arena, tuple)?; evaluator .as_ref() @@ -78,7 +88,7 @@ impl ScalarExpression { .binary_eval(&left, &right) } ScalarExpression::IsNull { expr, negated } => { - let mut is_null = expr.eval(tuple)?.is_null(); + let mut is_null = arena.expression(*expr).eval(arena, tuple)?.is_null(); if *negated { is_null = !is_null; } @@ -89,7 +99,7 @@ impl ScalarExpression { args, negated, } => { - let value = expr.eval(tuple)?; + let value = arena.expression(*expr).eval(arena, tuple)?; if value.is_null() { return Ok(DataValue::Null); } @@ -97,7 +107,7 @@ impl ScalarExpression { let mut matched = false; let mut saw_null = false; for arg in args { - let arg_value = arg.eval(tuple)?; + let arg_value = arena.expression(*arg).eval(arena, tuple)?; if arg_value.is_null() { saw_null = true; @@ -120,7 +130,7 @@ impl ScalarExpression { ScalarExpression::Unary { expr, evaluator, .. } => { - let value = expr.eval(tuple)?; + let value = arena.expression(*expr).eval(arena, tuple)?; Ok(evaluator .as_ref() @@ -136,9 +146,9 @@ impl ScalarExpression { right_expr, negated, } => { - let value = expr.eval(tuple)?; - let left = left_expr.eval(tuple)?; - let right = right_expr.eval(tuple)?; + let value = arena.expression(*expr).eval(arena, tuple)?; + let left = arena.expression(*left_expr).eval(arena, tuple)?; + let right = arena.expression(*right_expr).eval(arena, tuple)?; let mut is_between = match ( value.partial_cmp(&left).map(Ordering::is_ge), @@ -158,14 +168,15 @@ impl ScalarExpression { for_expr, from_expr, } => { - if let Some(mut string) = expr - .eval(tuple)? + if let Some(mut string) = arena + .expression(*expr) + .eval(arena, tuple)? .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? .utf8() .map(String::from) { if let Some(from_expr) = from_expr { - let mut from = eval_to_num!(from_expr, tuple).saturating_sub(1); + let mut from = eval_to_num!(from_expr, arena, tuple).saturating_sub(1); let len_i = string.len() as i32; while from < 0 { @@ -177,7 +188,8 @@ impl ScalarExpression { string = string.split_off(from as usize); } if let Some(for_expr) = for_expr { - let for_i = cmp::min(eval_to_num!(for_expr, tuple) as usize, string.len()); + let for_i = + cmp::min(eval_to_num!(for_expr, arena, tuple) as usize, string.len()); let _ = string.split_off(for_i); } @@ -191,16 +203,17 @@ impl ScalarExpression { } } ScalarExpression::Position { expr, in_expr } => { - let unpack = |expr: &ScalarExpression| -> Result { - Ok(expr - .eval(tuple)? + let unpack = |expr: ExprRef| -> Result { + Ok(arena + .expression(expr) + .eval(arena, tuple)? .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? .utf8() .map(String::from) .unwrap_or("".to_owned())) }; - let pattern = unpack(expr)?; - let str = unpack(in_expr)?; + let pattern = unpack(*expr)?; + let str = unpack(*in_expr)?; Ok(DataValue::Int32( str.find(&pattern).map(|pos| pos as i32 + 1).unwrap_or(0), )) @@ -210,15 +223,17 @@ impl ScalarExpression { trim_what_expr, trim_where, } => { - if let Some(string) = expr - .eval(tuple)? + if let Some(string) = arena + .expression(*expr) + .eval(arena, tuple)? .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? .utf8() { let mut trim_what = String::from(" "); if let Some(trim_what_expr) = trim_what_expr { - trim_what = trim_what_expr - .eval(tuple)? + trim_what = arena + .expression(*trim_what_expr) + .eval(arena, tuple)? .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? .utf8() .map(String::from) @@ -239,14 +254,14 @@ impl ScalarExpression { let mut values = Vec::with_capacity(exprs.len()); for expr in exprs { - values.push(expr.eval(tuple)?); + values.push(arena.expression(*expr).eval(arena, tuple)?); } Ok(DataValue::Tuple(values, false)) } ScalarExpression::ScalaFunction(ScalarFunction { inner, args, .. }) => { let value = match tuple { - Some(tuple) => inner.eval(args, Some(&tuple as &dyn TupleLike))?, - None => inner.eval(args, None)?, + Some(tuple) => inner.eval(args, arena, Some(&tuple as &dyn TupleLike))?, + None => inner.eval(args, arena, None)?, }; value.cast(inner.return_type()) } @@ -257,10 +272,10 @@ impl ScalarExpression { right_expr, ty, } => { - if condition.eval(tuple)?.is_true()? { - left_expr.eval(tuple)?.cast(ty) + if arena.expression(*condition).eval(arena, tuple)?.is_true()? { + arena.expression(*left_expr).eval(arena, tuple)?.cast(ty) } else { - right_expr.eval(tuple)?.cast(ty) + arena.expression(*right_expr).eval(arena, tuple)?.cast(ty) } } ScalarExpression::IfNull { @@ -268,10 +283,10 @@ impl ScalarExpression { right_expr, ty, } => { - let mut value = left_expr.eval(tuple)?; + let mut value = arena.expression(*left_expr).eval(arena, tuple)?; if value.is_null() { - value = right_expr.eval(tuple)?; + value = arena.expression(*right_expr).eval(arena, tuple)?; } value.cast(ty) } @@ -280,9 +295,9 @@ impl ScalarExpression { right_expr, ty, } => { - let mut value = left_expr.eval(tuple)?; + let mut value = arena.expression(*left_expr).eval(arena, tuple)?; - if right_expr.eval(tuple)? == value { + if arena.expression(*right_expr).eval(arena, tuple)? == value { value = DataValue::Null; } value.cast(ty) @@ -291,7 +306,7 @@ impl ScalarExpression { let mut value = None; for expr in exprs { - let temp = expr.eval(tuple)?; + let temp = arena.expression(*expr).eval(arena, tuple)?; if !temp.is_null() { value = Some(temp); @@ -310,10 +325,10 @@ impl ScalarExpression { let mut result = None; if let Some(expr) = operand_expr { - operand_value = Some(expr.eval(tuple)?); + operand_value = Some(arena.expression(*expr).eval(arena, tuple)?); } for (when_expr, result_expr) in expr_pairs { - let mut when_value = when_expr.eval(tuple)?; + let mut when_value = arena.expression(*when_expr).eval(arena, tuple)?; let is_true = if let Some(operand_value) = &operand_value { let ty = operand_value.logical_type(); when_value = when_value.cast(&ty)?; @@ -325,13 +340,13 @@ impl ScalarExpression { when_value.is_true()? }; if is_true { - result = Some(result_expr.eval(tuple)?); + result = Some(arena.expression(*result_expr).eval(arena, tuple)?); break; } } if result.is_none() { if let Some(expr) = else_expr { - result = Some(expr.eval(tuple)?); + result = Some(arena.expression(*expr).eval(arena, tuple)?); } } result.unwrap_or(DataValue::Null).cast(ty) @@ -373,47 +388,75 @@ fn trim_string(value: &str, trim_what: &str, trim_where: Option) mod tests { use super::*; - fn const_in(expr: DataValue, args: Vec, negated: bool) -> ScalarExpression { - ScalarExpression::In { + fn const_in( + arena: &mut PlanArena, + expr: DataValue, + args: Vec, + negated: bool, + ) -> ExprRef { + let expr = arena.alloc_expression(ScalarExpression::Constant(expr)); + let args = args + .into_iter() + .map(|value| arena.alloc_expression(ScalarExpression::Constant(value))) + .collect(); + arena.alloc_expression(ScalarExpression::In { negated, - expr: Box::new(ScalarExpression::Constant(expr)), - args: args.into_iter().map(ScalarExpression::Constant).collect(), - } + expr, + args, + }) } #[test] fn in_eval_matches_even_if_null_appears_first() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let expr = const_in( + &mut arena, DataValue::Int32(1), vec![DataValue::Null, DataValue::Int32(1)], false, ); - assert_eq!(expr.eval::<&[DataValue]>(None)?, DataValue::Boolean(true)); + assert_eq!( + arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + DataValue::Boolean(true) + ); Ok(()) } #[test] fn in_eval_returns_null_when_only_null_blocks_non_match() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let expr = const_in( + &mut arena, DataValue::Int32(2), vec![DataValue::Null, DataValue::Int32(1)], false, ); - assert_eq!(expr.eval::<&[DataValue]>(None)?, DataValue::Null); + assert_eq!( + arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + DataValue::Null + ); Ok(()) } #[test] fn not_in_eval_matches_even_if_null_appears_first() -> Result<(), DatabaseError> { + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let expr = const_in( + &mut arena, DataValue::Int32(1), vec![DataValue::Null, DataValue::Int32(1)], true, ); - assert_eq!(expr.eval::<&[DataValue]>(None)?, DataValue::Boolean(false)); + assert_eq!( + arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + DataValue::Boolean(false) + ); Ok(()) } diff --git a/src/expression/function/scala.rs b/src/expression/function/scala.rs index 74cac768..51d493a2 100644 --- a/src/expression/function/scala.rs +++ b/src/expression/function/scala.rs @@ -14,7 +14,7 @@ use crate::errors::DatabaseError; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, PlanArena}; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -43,7 +43,7 @@ impl Deref for ArcScalarFunctionImpl { #[derive(Debug, Clone, ReferenceSerialization)] pub struct ScalarFunction { - pub(crate) args: Vec, + pub(crate) args: Vec, pub(crate) inner: ArcScalarFunctionImpl, } @@ -63,7 +63,8 @@ impl Hash for ScalarFunction { pub trait ScalarFunctionImpl: Debug + Send + Sync { fn eval( &self, - args: &[ScalarExpression], + args: &[ExprRef], + arena: &PlanArena<'_>, tuple: Option<&dyn TupleLike>, ) -> Result; diff --git a/src/expression/function/table.rs b/src/expression/function/table.rs index b14a6b58..e01491ad 100644 --- a/src/expression/function/table.rs +++ b/src/expression/function/table.rs @@ -14,8 +14,7 @@ use crate::errors::DatabaseError; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; -use crate::planner::TableArena; +use crate::planner::{ExprRef, PlanArena, TableArena}; use crate::types::tuple::{Schema, Tuple}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt::Debug; @@ -36,7 +35,7 @@ impl Deref for ArcTableFunctionImpl { #[derive(Debug, Clone, ReferenceSerialization)] pub struct TableFunction { - pub(crate) args: Vec, + pub(crate) args: Vec, pub(crate) catalog: TableFunctionCatalog, } @@ -62,7 +61,8 @@ impl Hash for TableFunction { pub trait TableFunctionImpl: Debug + Send + Sync { fn eval( &self, - args: &[ScalarExpression], + args: &[ExprRef], + arena: &PlanArena<'_>, ) -> Result>>, DatabaseError>; fn summary(&self) -> &FunctionSummary; diff --git a/src/expression/mod.rs b/src/expression/mod.rs index eec778c2..29d5ab0d 100644 --- a/src/expression/mod.rs +++ b/src/expression/mod.rs @@ -19,9 +19,8 @@ use crate::expression::function::scala::ScalarFunction; use crate::expression::function::table::TableFunction; use crate::expression::visitor::{walk_expr, ExprVisitor}; use crate::expression::visitor_mut::ExprVisitorMut; -use crate::iter_ext::Itertools; use crate::planner::operator::sort::SortField; -use crate::planner::{MetaArena, PlanArena}; +use crate::planner::{Explain, ExprRef, MetaArena, PlanArena}; use crate::types::evaluator::{ binary_create, cast_create, unary_create, BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef, @@ -32,12 +31,13 @@ use kite_sql_serde_macros::ReferenceSerialization; #[cfg(feature = "decimal")] use rust_decimal::Decimal; use std::borrow::Cow; +use std::fmt; use std::fmt::{Debug, Formatter}; use std::hash::Hash; use std::sync::Arc; -use std::{fmt, mem}; pub mod agg; +mod eq_col; mod evaluator; pub mod function; pub mod range_detacher; @@ -56,7 +56,7 @@ pub enum TrimWhereField { #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub enum AliasType { Name(String), - Expr(Box), + Expr(ExprRef), } /// ScalarExpression represnet all scalar expression in SQL. @@ -71,91 +71,91 @@ pub enum ScalarExpression { position: usize, }, Alias { - expr: Box, + expr: ExprRef, alias: AliasType, }, TypeCast { - expr: Box, + expr: ExprRef, ty: LogicalType, evaluator: Option, }, IsNull { negated: bool, - expr: Box, + expr: ExprRef, }, Unary { op: UnaryOperator, - expr: Box, + expr: ExprRef, evaluator: Option, ty: LogicalType, }, Binary { op: BinaryOperator, - left_expr: Box, - right_expr: Box, + left_expr: ExprRef, + right_expr: ExprRef, evaluator: Option, ty: LogicalType, }, AggCall { distinct: bool, kind: AggKind, - args: Vec, + args: Vec, ty: LogicalType, }, In { negated: bool, - expr: Box, - args: Vec, + expr: ExprRef, + args: Vec, }, Between { negated: bool, - expr: Box, - left_expr: Box, - right_expr: Box, + expr: ExprRef, + left_expr: ExprRef, + right_expr: ExprRef, }, SubString { - expr: Box, - for_expr: Option>, - from_expr: Option>, + expr: ExprRef, + for_expr: Option, + from_expr: Option, }, Position { - expr: Box, - in_expr: Box, + expr: ExprRef, + in_expr: ExprRef, }, Trim { - expr: Box, - trim_what_expr: Option>, + expr: ExprRef, + trim_what_expr: Option, trim_where: Option, }, // Temporary expression used for expression substitution Empty, - Tuple(Vec), + Tuple(Vec), ScalaFunction(ScalarFunction), TableFunction(TableFunction), If { - condition: Box, - left_expr: Box, - right_expr: Box, + condition: ExprRef, + left_expr: ExprRef, + right_expr: ExprRef, ty: LogicalType, }, IfNull { - left_expr: Box, - right_expr: Box, + left_expr: ExprRef, + right_expr: ExprRef, ty: LogicalType, }, NullIf { - left_expr: Box, - right_expr: Box, + left_expr: ExprRef, + right_expr: ExprRef, ty: LogicalType, }, Coalesce { - exprs: Vec, + exprs: Vec, ty: LogicalType, }, CaseWhen { - operand_expr: Option>, - expr_pairs: Vec<(ScalarExpression, ScalarExpression)>, - else_expr: Option>, + operand_expr: Option, + expr_pairs: Vec<(ExprRef, ExprRef)>, + else_expr: Option, ty: LogicalType, }, WindowCall(window::WindowCall), @@ -275,23 +275,22 @@ mod chrono_scalar_expression { } } -pub struct BindEvaluator<'a, 'p> { - pub(crate) arena: &'a PlanArena<'p>, -} +pub struct BindEvaluator; -impl ExprVisitorMut<'_> for BindEvaluator<'_, '_> { +impl ExprVisitorMut for BindEvaluator { fn visit_type_cast( &mut self, - expr: &'_ mut ScalarExpression, - ty: &'_ mut LogicalType, - evaluator: &'_ mut Option, + expr: &mut ExprRef, + ty: &mut LogicalType, + evaluator: &mut Option, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - let from = expr.return_type(self.arena); + self.visit(expr, arena)?; + let from = expr.return_type(arena); *evaluator = if from.as_ref() == ty { None } else { - Some(cast_create(from, Cow::Borrowed(ty))?) + Some(cast_create(from.as_ref(), ty)?) }; Ok(()) @@ -300,13 +299,14 @@ impl ExprVisitorMut<'_> for BindEvaluator<'_, '_> { fn visit_unary( &mut self, op: &'_ mut UnaryOperator, - expr: &'_ mut ScalarExpression, - evaluator: &'_ mut Option, - _ty: &'_ mut LogicalType, + expr: &mut ExprRef, + evaluator: &mut Option, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; + self.visit(expr, arena)?; - let ty = expr.return_type(self.arena); + let ty = expr.return_type(arena); if ty.is_unsigned_numeric() { let target_ty = match ty.as_ref() { LogicalType::UTinyint => LogicalType::Tinyint, @@ -315,13 +315,9 @@ impl ExprVisitorMut<'_> for BindEvaluator<'_, '_> { LogicalType::UBigint => LogicalType::Bigint, _ => unreachable!(), }; - *expr = ScalarExpression::type_cast( - mem::replace(expr, ScalarExpression::Empty), - Cow::Owned(target_ty), - self.arena, - )?; + *expr = (*expr).type_cast(Cow::Owned(target_ty), arena)?; } - *evaluator = Some(unary_create(expr.return_type(self.arena), *op)?); + *evaluator = Some(unary_create(expr.return_type(arena), *op)?); Ok(()) } @@ -329,30 +325,22 @@ impl ExprVisitorMut<'_> for BindEvaluator<'_, '_> { fn visit_binary( &mut self, op: &'_ mut BinaryOperator, - left_expr: &'_ mut ScalarExpression, - right_expr: &'_ mut ScalarExpression, - evaluator: &'_ mut Option, - _ty: &'_ mut LogicalType, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + evaluator: &mut Option, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr)?; - - let left_ty = left_expr.return_type(self.arena).into_owned(); - let right_ty = right_expr.return_type(self.arena).into_owned(); - let ty = LogicalType::max_logical_type(&left_ty, &right_ty)?; - let fn_cast = - |expr: &mut ScalarExpression, ty: &LogicalType| -> Result<(), DatabaseError> { - *expr = ScalarExpression::type_cast( - mem::replace(expr, ScalarExpression::Empty), - Cow::Borrowed(ty), - self.arena, - )?; - Ok(()) - }; - fn_cast(left_expr, ty.as_ref())?; - fn_cast(right_expr, ty.as_ref())?; + self.visit(left_expr, arena)?; + self.visit(right_expr, arena)?; + + let left_ty = left_expr.return_type(arena).into_owned(); + let right_ty = right_expr.return_type(arena).into_owned(); + let ty = LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned(); + *left_expr = left_expr.type_cast(Cow::Borrowed(&ty), arena)?; + *right_expr = right_expr.type_cast(Cow::Borrowed(&ty), arena)?; - *evaluator = Some(binary_create(ty, *op)?); + *evaluator = Some(binary_create(Cow::Owned(ty), *op)?); Ok(()) } @@ -363,230 +351,425 @@ pub struct HasCountStar { pub value: bool, } -impl ExprVisitor<'_> for HasCountStar { +impl ExprVisitor> for HasCountStar { fn visit_agg( &mut self, _distinct: bool, - _kind: &'_ AggKind, - args: &'_ [ScalarExpression], - _ty: &'_ LogicalType, + _kind: &AggKind, + args: &[ExprRef], + _ty: &LogicalType, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { if args.len() == 1 { - if let ScalarExpression::Constant(value) = &args[0] { + if let ScalarExpression::Constant(value) = arena.expression(args[0]) { self.value = matches!(value.utf8(), Some("*")); } } Ok(()) } - fn visit(&mut self, expr: &'_ ScalarExpression) -> Result<(), DatabaseError> { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { if !self.value { - walk_expr(self, expr)?; + walk_expr(self, expr, arena)?; } Ok(()) } } -impl ScalarExpression { - pub fn asc(self) -> SortField { - SortField::from(self).asc() - } +pub trait TypeCast: Sized { + fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType>; - pub fn desc(self) -> SortField { - SortField::from(self).desc() - } - - pub fn nulls_first(self) -> SortField { - SortField::from(self).nulls_first() - } - - pub fn nulls_last(self) -> SortField { - SortField::from(self).nulls_last() - } - - pub fn column_expr(column: ColumnRef, position: usize) -> ScalarExpression { - ScalarExpression::ColumnRef { column, position } - } + fn into_expr( + self, + ty: LogicalType, + evaluator: CastEvaluatorRef, + arena: &mut PlanArena<'_>, + ) -> Self; - pub fn type_cast( - expr: ScalarExpression, + fn type_cast( + self, ty: Cow<'_, LogicalType>, - arena: &PlanArena, - ) -> Result { - let from = expr.return_type(arena); + arena: &mut PlanArena<'_>, + ) -> Result { + let from = self.return_type(arena); if from.as_ref() == ty.as_ref() { - return Ok(expr); - } - let evaluator = Some(cast_create(from, ty.clone())?); - - Ok(ScalarExpression::TypeCast { - expr: Box::new(expr), - ty: ty.into_owned(), - evaluator, - }) - } - - pub(crate) fn eq_ignore_colref_pos(&self, other: &ScalarExpression, arena: &PlanArena) -> bool { - match (self.unpack_alias_ref(), other.unpack_alias_ref()) { - ( - ScalarExpression::ColumnRef { - column: lhs_column, .. - }, - ScalarExpression::ColumnRef { - column: rhs_column, .. - }, - ) => arena.same_column(*lhs_column, *rhs_column), - (lhs, rhs) => lhs == rhs, - } - } - - pub fn unpack_alias(self) -> ScalarExpression { - if let ScalarExpression::Alias { - alias: AliasType::Expr(expr), - .. - } = self - { - expr.unpack_alias() - } else if let ScalarExpression::Alias { expr, .. } = self { - expr.unpack_alias() - } else { - self - } - } - - pub fn unpack_alias_ref(&self) -> &ScalarExpression { - if let ScalarExpression::Alias { - alias: AliasType::Expr(expr), - .. - } = self - { - expr.unpack_alias_ref() - } else if let ScalarExpression::Alias { expr, .. } = self { - expr.unpack_alias_ref() - } else { - self + return Ok(self); } + let evaluator = cast_create(from.as_ref(), ty.as_ref())?; + Ok(self.into_expr(ty.into_owned(), evaluator, arena)) } +} - pub fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType> { +impl TypeCast for ScalarExpression { + fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType> { match self { - ScalarExpression::Constant(v) => Cow::Owned(v.logical_type()), + ScalarExpression::Constant(value) => Cow::Owned(value.logical_type()), ScalarExpression::ColumnRef { column, .. } => { Cow::Borrowed(arena.column(*column).datatype()) } - ScalarExpression::Binary { - ty: return_type, .. - } - | ScalarExpression::Unary { - ty: return_type, .. - } - | ScalarExpression::TypeCast { - ty: return_type, .. - } - | ScalarExpression::AggCall { - ty: return_type, .. - } - | ScalarExpression::If { - ty: return_type, .. - } - | ScalarExpression::IfNull { - ty: return_type, .. - } - | ScalarExpression::NullIf { - ty: return_type, .. - } - | ScalarExpression::Coalesce { - ty: return_type, .. - } - | ScalarExpression::CaseWhen { - ty: return_type, .. - } + ScalarExpression::Binary { ty, .. } + | ScalarExpression::Unary { ty, .. } + | ScalarExpression::TypeCast { ty, .. } + | ScalarExpression::AggCall { ty, .. } + | ScalarExpression::If { ty, .. } + | ScalarExpression::IfNull { ty, .. } + | ScalarExpression::NullIf { ty, .. } + | ScalarExpression::Coalesce { ty, .. } + | ScalarExpression::CaseWhen { ty, .. } | ScalarExpression::WindowCall(window::WindowCall { - function: - window::WindowFunction { - ty: return_type, .. - }, + function: window::WindowFunction { ty, .. }, .. - }) => Cow::Borrowed(return_type), + }) => Cow::Borrowed(ty), ScalarExpression::IsNull { .. } | ScalarExpression::In { .. } | ScalarExpression::Between { .. } => Cow::Owned(LogicalType::Boolean), - ScalarExpression::SubString { .. } => { + ScalarExpression::SubString { .. } | ScalarExpression::Trim { .. } => { Cow::Owned(LogicalType::Varchar(None, CharLengthUnits::Characters)) } ScalarExpression::Position { .. } => Cow::Owned(LogicalType::Integer), - ScalarExpression::Trim { .. } => { - Cow::Owned(LogicalType::Varchar(None, CharLengthUnits::Characters)) - } ScalarExpression::Alias { expr, .. } => expr.return_type(arena), ScalarExpression::Empty | ScalarExpression::TableFunction(_) => unreachable!(), - ScalarExpression::Tuple(exprs) => { - let types = exprs + ScalarExpression::Tuple(exprs) => Cow::Owned(LogicalType::Tuple( + exprs .iter() .map(|expr| expr.return_type(arena).into_owned()) - .collect_vec(); - - Cow::Owned(LogicalType::Tuple(types)) - } + .collect(), + )), ScalarExpression::ScalaFunction(ScalarFunction { inner, .. }) => { Cow::Borrowed(inner.return_type()) } } } - pub fn visit_referenced_columns( - &self, - arena: &mut A, - f: &mut impl FnMut(&mut A, &ColumnRef) -> bool, - ) -> Result { - struct ColumnRefVisitor<'a, A, F> { - f: &'a mut F, - keep_going: bool, - arena: &'a mut A, + fn into_expr( + self, + ty: LogicalType, + evaluator: CastEvaluatorRef, + arena: &mut PlanArena<'_>, + ) -> Self { + ScalarExpression::TypeCast { + expr: arena.alloc_expression(self), + ty, + evaluator: Some(evaluator), } + } +} - impl ExprVisitor<'_> for ColumnRefVisitor<'_, A, F> - where - A: MetaArena, - F: FnMut(&mut A, &ColumnRef) -> bool, - { - fn visit(&mut self, expr: &ScalarExpression) -> Result<(), DatabaseError> { - if self.keep_going { - walk_expr(self, expr)?; +impl TypeCast for ExprRef { + fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType> { + arena.expression(*self).return_type(arena) + } + + fn into_expr( + self, + ty: LogicalType, + evaluator: CastEvaluatorRef, + arena: &mut PlanArena<'_>, + ) -> Self { + arena.alloc_expression(ScalarExpression::TypeCast { + expr: self, + ty, + evaluator: Some(evaluator), + }) + } +} + +impl ScalarExpression { + pub fn column_expr(column: ColumnRef, position: usize) -> ScalarExpression { + ScalarExpression::ColumnRef { column, position } + } +} + +impl Explain for ExprRef { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn write_exprs( + exprs: &[ExprRef], + arena: &PlanArena<'_>, + f: &mut fmt::Formatter<'_>, + ) -> fmt::Result { + for (index, expr) in exprs.iter().enumerate() { + if index > 0 { + f.write_str(", ")?; } - Ok(()) + write!(f, "{}", expr.explain(arena))?; } + Ok(()) + } - fn visit_column_ref(&mut self, col: &ColumnRef) -> Result<(), DatabaseError> { - self.keep_going = (self.f)(self.arena, col); - Ok(()) + match arena.expression(*self) { + ScalarExpression::Constant(value) => write!(f, "{value}"), + ScalarExpression::ColumnRef { column, .. } => Explain::fmt(column, arena, f), + ScalarExpression::Alias { alias, expr } => match alias { + AliasType::Name(alias) => f.write_str(alias), + AliasType::Expr(alias_expr) => write!( + f, + "({}) as ({})", + expr.explain(arena), + alias_expr.explain(arena) + ), + }, + ScalarExpression::TypeCast { expr, ty, .. } => { + write!(f, "cast ({} as {ty})", expr.explain(arena)) + } + ScalarExpression::IsNull { expr, negated } => write!( + f, + "{} {}", + expr.explain(arena), + if *negated { "is not null" } else { "is null" } + ), + ScalarExpression::Unary { expr, op, .. } => { + write!(f, "{}{}", op, expr.explain(arena)) + } + ScalarExpression::Binary { + left_expr, + right_expr, + op, + .. + } => write!( + f, + "({} {op} {})", + left_expr.explain(arena), + right_expr.explain(arena) + ), + ScalarExpression::AggCall { + args, + kind, + distinct, + .. + } => { + write!(f, "{kind:?}(")?; + if kind.allow_distinct() && *distinct { + f.write_str("distinct ")?; + } + write_exprs(args, arena, f)?; + f.write_str(")") + } + ScalarExpression::WindowCall(window) => { + write!(f, "{}(", window.function.kind.name())?; + write_exprs(&window.function.args, arena, f)?; + f.write_str(") over (")?; + let mut has_spec = false; + if !window.spec.partition_by.is_empty() { + f.write_str("partition by ")?; + write_exprs(&window.spec.partition_by, arena, f)?; + has_spec = true; + } + if !window.spec.order_by.is_empty() { + if has_spec { + f.write_str(" ")?; + } + f.write_str("order by ")?; + for (index, field) in window.spec.order_by.iter().enumerate() { + if index > 0 { + f.write_str(", ")?; + } + write!(f, "{}", field.explain(arena))?; + } + } + f.write_str(")") + } + ScalarExpression::In { + args, + negated, + expr, + } => { + write!( + f, + "{} {} (", + expr.explain(arena), + if *negated { "not in" } else { "in" } + )?; + write_exprs(args, arena, f)?; + f.write_str(")") + } + ScalarExpression::Between { + expr, + left_expr, + right_expr, + negated, + } => write!( + f, + "{} {} [{}, {}]", + expr.explain(arena), + if *negated { "not between" } else { "between" }, + left_expr.explain(arena), + right_expr.explain(arena) + ), + ScalarExpression::SubString { + expr, + for_expr, + from_expr, + } => { + write!(f, "substring({}", expr.explain(arena))?; + if let Some(from_expr) = from_expr { + write!(f, ", from: {}", from_expr.explain(arena))?; + } + if let Some(for_expr) = for_expr { + write!(f, ", for: {}", for_expr.explain(arena))?; + } + f.write_str(")") + } + ScalarExpression::Position { expr, in_expr } => write!( + f, + "position({} in {})", + expr.explain(arena), + in_expr.explain(arena) + ), + ScalarExpression::Trim { + expr, + trim_what_expr, + trim_where, + } => { + let trim_what = trim_what_expr + .as_ref() + .map(|expr| expr.explain(arena).to_string()) + .unwrap_or_else(|| " ".to_string()); + + f.write_str("trim(")?; + match trim_where { + Some(TrimWhereField::Both) => write!(f, "both '{trim_what}' from")?, + Some(TrimWhereField::Leading) => write!(f, "leading '{trim_what}' from")?, + Some(TrimWhereField::Trailing) => write!(f, "trailing '{trim_what}' from")?, + None if !trim_what.is_empty() => write!(f, "'{trim_what}' from")?, + None => {} + } + write!(f, " {})", expr.explain(arena)) + } + ScalarExpression::Empty => unreachable!(), + ScalarExpression::Tuple(args) => { + f.write_str("(")?; + write_exprs(args, arena, f)?; + f.write_str(")") + } + ScalarExpression::ScalaFunction(ScalarFunction { args, inner }) => { + write!(f, "{}(", inner.summary().name)?; + write_exprs(args, arena, f)?; + f.write_str(")") + } + ScalarExpression::TableFunction(TableFunction { args, catalog }) => { + write!(f, "{}(", catalog.inner.summary().name)?; + write_exprs(args, arena, f)?; + f.write_str(")") + } + ScalarExpression::If { + condition, + left_expr, + right_expr, + .. + } => write!( + f, + "if {} ({}, {})", + condition.explain(arena), + left_expr.explain(arena), + right_expr.explain(arena) + ), + ScalarExpression::IfNull { + left_expr, + right_expr, + .. + } + | ScalarExpression::NullIf { + left_expr, + right_expr, + .. + } => write!( + f, + "ifnull({}, {})", + left_expr.explain(arena), + right_expr.explain(arena) + ), + ScalarExpression::Coalesce { exprs, .. } => { + f.write_str("coalesce(")?; + write_exprs(exprs, arena, f)?; + f.write_str(")") + } + ScalarExpression::CaseWhen { + operand_expr, + expr_pairs, + else_expr, + .. + } => { + f.write_str("case ")?; + if let Some(operand_expr) = operand_expr { + write!(f, "{} ", operand_expr.explain(arena))?; + } + for (index, (when_expr, then_expr)) in expr_pairs.iter().enumerate() { + if index > 0 { + f.write_str(" ")?; + } + write!( + f, + "when {} then {}", + when_expr.explain(arena), + then_expr.explain(arena) + )?; + } + f.write_str(" ")?; + if let Some(else_expr) = else_expr { + write!(f, "else {} ", else_expr.explain(arena))?; + } + f.write_str("end") } } + } +} - let mut visitor = ColumnRefVisitor { - f, - keep_going: true, - arena, - }; - visitor.visit(self)?; - Ok(visitor.keep_going) +impl ExprRef { + pub fn asc(self) -> SortField { + SortField::from(self).asc() + } + + pub fn desc(self) -> SortField { + SortField::from(self).desc() + } + + pub fn nulls_first(self) -> SortField { + SortField::from(self).nulls_first() + } + + pub fn nulls_last(self) -> SortField { + SortField::from(self).nulls_last() + } + + pub(crate) fn eq_ignore_colref_pos(self, other: ExprRef, arena: &PlanArena) -> bool { + eq_col::eq_ignore_colref_pos(self, other, arena) + } + + pub fn unpack_alias(self, arena: &impl MetaArena) -> ExprRef { + if let ScalarExpression::Alias { + alias: AliasType::Expr(expr), + .. + } = arena.expression(self) + { + expr.unpack_alias(arena) + } else if let ScalarExpression::Alias { expr, .. } = arena.expression(self) { + expr.unpack_alias(arena) + } else { + self + } + } + + pub fn unpack_alias_ref<'a, A: MetaArena>(self, arena: &'a A) -> &'a ScalarExpression { + arena.expression(self.unpack_alias(arena)) } pub fn any_referenced_column( - &self, + self, arena: &PlanArena, mut predicate: impl FnMut(&PlanArena, &ColumnRef) -> bool, ) -> Result { - struct ColumnRefVisitor<'a, 'p, F> { + struct ColumnRefVisitor<'a, 'arena, F> { f: &'a mut F, any: bool, - arena: &'a PlanArena<'p>, + arena: &'a PlanArena<'arena>, } - impl bool> ExprVisitor<'_> for ColumnRefVisitor<'_, '_, F> { - fn visit(&mut self, expr: &ScalarExpression) -> Result<(), DatabaseError> { + impl bool> ExprVisitor> + for ColumnRefVisitor<'_, '_, F> + { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { if !self.any { - walk_expr(self, expr)?; + walk_expr(self, expr, arena)?; } Ok(()) } @@ -602,25 +785,27 @@ impl ScalarExpression { any: false, arena, }; - visitor.visit(self)?; + visitor.visit(self, arena)?; Ok(visitor.any) } pub fn all_referenced_columns( - &self, + self, arena: &PlanArena, mut predicate: impl FnMut(&PlanArena, &ColumnRef) -> bool, ) -> Result { - struct ColumnRefVisitor<'a, 'p, F> { + struct ColumnRefVisitor<'a, 'arena, F> { f: &'a mut F, all: bool, - arena: &'a PlanArena<'p>, + arena: &'a PlanArena<'arena>, } - impl bool> ExprVisitor<'_> for ColumnRefVisitor<'_, '_, F> { - fn visit(&mut self, expr: &ScalarExpression) -> Result<(), DatabaseError> { + impl bool> ExprVisitor> + for ColumnRefVisitor<'_, '_, F> + { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { if self.all { - walk_expr(self, expr)?; + walk_expr(self, expr, arena)?; } Ok(()) } @@ -636,76 +821,56 @@ impl ScalarExpression { all: true, arena, }; - visitor.visit(self)?; + visitor.visit(self, arena)?; Ok(visitor.all) } - pub fn has_table_ref_column(&self, arena: &PlanArena) -> Result { - struct TableRefChecker<'arena, 'table> { - found: bool, - arena: &'arena PlanArena<'table>, - } - impl ExprVisitor<'_> for TableRefChecker<'_, '_> { - fn visit_column_ref(&mut self, col: &ColumnRef) -> Result<(), DatabaseError> { - let col = self.arena.column(*col); - if col.table_name().is_some() && col.id().is_some() { - self.found = true; - } - Ok(()) - } - } - let mut checker = TableRefChecker { - found: false, - arena, - }; - checker.visit(self)?; - Ok(checker.found) - } - - pub fn has_agg_call(&self) -> Result { + pub fn has_agg_call(self, arena: &PlanArena<'_>) -> Result { struct AggCallChecker { has_agg: bool, } - impl<'a> ExprVisitor<'a> for AggCallChecker { - fn visit(&mut self, expr: &'a ScalarExpression) -> Result<(), DatabaseError> { + impl ExprVisitor> for AggCallChecker { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { if self.has_agg { return Ok(()); } - walk_expr(self, expr) + walk_expr(self, expr, arena) } fn visit_agg( &mut self, _distinct: bool, - _kind: &'a AggKind, - args: &'a [ScalarExpression], - _ty: &'a LogicalType, + _kind: &AggKind, + args: &[ExprRef], + _ty: &LogicalType, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { for arg in args { - self.visit(arg)?; + self.visit(*arg, arena)?; } self.has_agg = true; Ok(()) } } let mut checker = AggCallChecker { has_agg: false }; - checker.visit(self)?; + checker.visit(self, arena)?; Ok(checker.has_agg) } - pub fn has_window_call(&self) -> Result { + pub fn has_window_call(self, arena: &PlanArena<'_>) -> Result { struct WindowCallChecker(bool); - impl<'a> ExprVisitor<'a> for WindowCallChecker { - fn visit(&mut self, expr: &'a ScalarExpression) -> Result<(), DatabaseError> { + impl ExprVisitor> for WindowCallChecker { + fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { if !self.0 { - walk_expr(self, expr)?; + walk_expr(self, expr, arena)?; } Ok(()) } fn visit_window( &mut self, - _window: &'a window::WindowCall, + _window: &window::WindowCall, + _arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { self.0 = true; Ok(()) @@ -713,292 +878,16 @@ impl ScalarExpression { } let mut checker = WindowCallChecker(false); - checker.visit(self)?; + checker.visit(self, arena)?; Ok(checker.0) } - fn output_name_by(&self, fn_display: &impl Fn(ColumnRef) -> N) -> String { - match self { - ScalarExpression::Constant(value) => format!("{value}"), - ScalarExpression::ColumnRef { column, .. } => format!("{}", fn_display(*column)), - ScalarExpression::Alias { alias, expr } => match alias { - AliasType::Name(alias) => alias.to_string(), - AliasType::Expr(alias_expr) => { - format!( - "({}) as ({})", - expr.output_name_by(fn_display), - alias_expr.output_name_by(fn_display) - ) - } - }, - ScalarExpression::TypeCast { expr, ty, .. } => { - format!("cast ({} as {})", expr.output_name_by(fn_display), ty) - } - ScalarExpression::IsNull { expr, negated } => { - let suffix = if *negated { "is not null" } else { "is null" }; - - format!("{} {}", expr.output_name_by(fn_display), suffix) - } - ScalarExpression::Unary { expr, op, .. } => { - format!("{}{}", op, expr.output_name_by(fn_display)) - } - ScalarExpression::Binary { - left_expr, - right_expr, - op, - .. - } => format!( - "({} {} {})", - left_expr.output_name_by(fn_display), - op, - right_expr.output_name_by(fn_display), - ), - ScalarExpression::AggCall { - args, - kind, - distinct, - .. - } => { - let args_str = args - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - let op = |allow_distinct, distinct| { - if allow_distinct && distinct { - "distinct " - } else { - "" - } - }; - format!( - "{:?}({}{})", - kind, - op(kind.allow_distinct(), *distinct), - args_str - ) - } - ScalarExpression::WindowCall(window) => { - let args = window - .function - .args - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - let function = window.function.kind.name(); - let mut spec = Vec::new(); - if !window.spec.partition_by.is_empty() { - spec.push(format!( - "partition by {}", - window - .spec - .partition_by - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", ") - )); - } - if !window.spec.order_by.is_empty() { - spec.push(format!( - "order by {}", - window - .spec - .order_by - .iter() - .map(ToString::to_string) - .join(", ") - )); - } - format!("{function}({args}) over ({})", spec.join(" ")) - } - ScalarExpression::In { - args, - negated, - expr, - } => { - let args_string = args - .iter() - .map(|arg| arg.output_name_by(fn_display)) - .join(", "); - let op_string = if *negated { "not in" } else { "in" }; - format!( - "{} {} ({})", - expr.output_name_by(fn_display), - op_string, - args_string - ) - } - ScalarExpression::Between { - expr, - left_expr, - right_expr, - negated, - } => { - let op_string = if *negated { "not between" } else { "between" }; - format!( - "{} {} [{}, {}]", - expr.output_name_by(fn_display), - op_string, - left_expr.output_name_by(fn_display), - right_expr.output_name_by(fn_display) - ) - } - ScalarExpression::SubString { - expr, - for_expr, - from_expr, - } => { - let op = |tag: &str, num_expr: &Option>| { - num_expr - .as_ref() - .map(|expr| format!(", {}: {}", tag, expr.output_name_by(fn_display))) - .unwrap_or_default() - }; - - format!( - "substring({}{}{})", - expr.output_name_by(fn_display), - op("from", from_expr), - op("for", for_expr), - ) - } - ScalarExpression::Position { expr, in_expr } => { - format!( - "position({} in {})", - expr.output_name_by(fn_display), - in_expr.output_name_by(fn_display) - ) - } - ScalarExpression::Trim { - expr, - trim_what_expr, - trim_where, - } => { - let trim_what_str = { - trim_what_expr - .as_ref() - .map(|expr| expr.output_name_by(fn_display)) - .unwrap_or_else(|| " ".to_string()) - }; - let trim_where_str = match trim_where { - Some(TrimWhereField::Both) => format!("both '{trim_what_str}' from"), - Some(TrimWhereField::Leading) => format!("leading '{trim_what_str}' from"), - Some(TrimWhereField::Trailing) => format!("trailing '{trim_what_str}' from"), - None => { - if trim_what_str.is_empty() { - String::new() - } else { - format!("'{trim_what_str}' from") - } - } - }; - format!( - "trim({} {})", - trim_where_str, - expr.output_name_by(fn_display) - ) - } - ScalarExpression::Empty => unreachable!(), - ScalarExpression::Tuple(args) => { - let args_str = args - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - format!("({args_str})") - } - ScalarExpression::ScalaFunction(ScalarFunction { args, inner }) => { - let args_str = args - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - format!("{}({})", inner.summary().name, args_str) - } - ScalarExpression::TableFunction(TableFunction { args, catalog }) => { - let args_str = args - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - format!("{}({})", catalog.inner.summary().name, args_str) - } - ScalarExpression::If { - condition, - left_expr, - right_expr, - .. - } => { - format!( - "if {} ({}, {})", - condition.output_name_by(fn_display), - left_expr.output_name_by(fn_display), - right_expr.output_name_by(fn_display) - ) - } - ScalarExpression::IfNull { - left_expr, - right_expr, - .. - } => { - format!( - "ifnull({}, {})", - left_expr.output_name_by(fn_display), - right_expr.output_name_by(fn_display) - ) - } - ScalarExpression::NullIf { - left_expr, - right_expr, - .. - } => { - format!( - "ifnull({}, {})", - left_expr.output_name_by(fn_display), - right_expr.output_name_by(fn_display) - ) - } - ScalarExpression::Coalesce { exprs, .. } => { - let exprs_str = exprs - .iter() - .map(|expr| expr.output_name_by(fn_display)) - .join(", "); - format!("coalesce({exprs_str})") - } - ScalarExpression::CaseWhen { - operand_expr, - expr_pairs, - else_expr, - .. - } => { - let op = |tag: &str, expr: &Option>| { - expr.as_ref() - .map(|expr| format!("{}{} ", tag, expr.output_name_by(fn_display))) - .unwrap_or_default() - }; - let expr_pairs_str = expr_pairs - .iter() - .map(|(when_expr, then_expr)| { - format!( - "when {} then {}", - when_expr.output_name_by(fn_display), - then_expr.output_name_by(fn_display) - ) - }) - .join(" "); - - format!( - "case {}{} {}end", - op("", operand_expr), - expr_pairs_str, - op("else ", else_expr) - ) - } - } - } - - pub fn output_name(&self, arena: &PlanArena) -> String { - self.output_name_by(&|column| arena.column(column).full_name()) + pub fn output_name(self, arena: &PlanArena) -> String { + self.explain(arena).to_string() } - pub fn output_column_ref(&self, arena: &mut PlanArena) -> ColumnRef { - match self { + pub fn output_column_ref(self, arena: &mut PlanArena) -> ColumnRef { + match arena.expression(self) { ScalarExpression::ColumnRef { column, .. } => *column, ScalarExpression::Alias { alias: AliasType::Expr(expr), @@ -1050,12 +939,6 @@ pub enum BinaryOperator { Or, } -impl fmt::Display for ScalarExpression { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - write!(f, "{}", self.output_name_by(&|column| column)) - } -} - impl fmt::Display for BinaryOperator { fn fmt(&self, f: &mut Formatter) -> fmt::Result { let like_op = |f: &mut Formatter, escape_char: &Option| { @@ -1121,7 +1004,7 @@ mod test { use crate::expression::{AliasType, BinaryOperator, ScalarExpression, UnaryOperator}; use crate::function::current_date::CurrentDate; use crate::function::numbers::Numbers; - use crate::planner::{PlanArena, TableArenaCell}; + use crate::planner::{ExprRef, PlanArena, TableArenaCell}; use crate::serdes::{ReferenceDecodeContext, ReferenceSerialization, ReferenceTables}; use crate::storage::rocksdb::RocksStorage; use crate::storage::rocksdb::RocksTransaction; @@ -1134,40 +1017,6 @@ mod test { use std::io::{Cursor, Seek, SeekFrom}; use tempfile::TempDir; - #[test] - fn test_eq_ignore_colref_pos() -> Result<(), DatabaseError> { - let table_arena = TableArenaCell::default(); - let mut arena = PlanArena::new(&table_arena); - let left = ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new( - "c1".to_string(), - false, - ColumnDesc::new(LogicalType::Integer, None, false, None)?, - )), - 0, - ); - let right = ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new( - "c1".to_string(), - true, - ColumnDesc::new(LogicalType::Bigint, None, false, None)?, - )), - 2, - ); - let different = ScalarExpression::column_expr( - arena.alloc_column(ColumnCatalog::new( - "c2".to_string(), - false, - ColumnDesc::new(LogicalType::Integer, None, false, None)?, - )), - 0, - ); - - assert!(left.eq_ignore_colref_pos(&right, &arena)); - assert!(!left.eq_ignore_colref_pos(&different, &arena)); - Ok(()) - } - #[test] fn test_serialization() -> Result<(), DatabaseError> { fn fn_assert( @@ -1177,12 +1026,13 @@ mod test { reference_tables: &mut ReferenceTables, arena: &mut PlanArena, ) -> Result<(), DatabaseError> { + let expr = arena.alloc_expression(expr); expr.encode(cursor, false, reference_tables, arena)?; cursor.seek(SeekFrom::Start(0))?; - let decoded = ScalarExpression::decode(cursor, drive, reference_tables, arena)?; + let decoded = ExprRef::decode(cursor, drive, reference_tables, arena)?; assert!( - decoded.eq_ignore_colref_pos(&expr, arena), + decoded.eq_ignore_colref_pos(expr, arena), "decoded expression does not match: decoded={decoded:?}, expected={expr:?}", ); cursor.seek(SeekFrom::Start(0))?; @@ -1226,6 +1076,7 @@ mod test { &scala_functions, &table_functions, ); + let empty = plan_arena.alloc_expression(ScalarExpression::Empty); fn_assert( &mut cursor, @@ -1276,7 +1127,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Alias { - expr: Box::new(ScalarExpression::Empty), + expr: empty, alias: AliasType::Name("Hello".to_string()), }, Some(&context), @@ -1286,8 +1137,8 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Alias { - expr: Box::new(ScalarExpression::Empty), - alias: AliasType::Expr(Box::new(ScalarExpression::Empty)), + expr: empty, + alias: AliasType::Expr(empty), }, Some(&context), &mut reference_tables, @@ -1296,12 +1147,9 @@ mod test { fn_assert( &mut cursor, ScalarExpression::TypeCast { - expr: Box::new(ScalarExpression::Empty), + expr: empty, ty: LogicalType::Integer, - evaluator: Some(cast_create( - Cow::Owned(LogicalType::Integer), - Cow::Owned(LogicalType::Integer), - )?), + evaluator: Some(cast_create(&LogicalType::Integer, &LogicalType::Integer)?), }, Some(&context), &mut reference_tables, @@ -1311,7 +1159,7 @@ mod test { &mut cursor, ScalarExpression::IsNull { negated: true, - expr: Box::new(ScalarExpression::Empty), + expr: empty, }, Some(&context), &mut reference_tables, @@ -1321,7 +1169,7 @@ mod test { &mut cursor, ScalarExpression::Unary { op: UnaryOperator::Plus, - expr: Box::new(ScalarExpression::Empty), + expr: empty, evaluator: Some(unary_create( Cow::Owned(LogicalType::Boolean), UnaryOperator::Not, @@ -1336,7 +1184,7 @@ mod test { &mut cursor, ScalarExpression::Unary { op: UnaryOperator::Plus, - expr: Box::new(ScalarExpression::Empty), + expr: empty, evaluator: None, ty: LogicalType::Integer, }, @@ -1348,8 +1196,8 @@ mod test { &mut cursor, ScalarExpression::Binary { op: BinaryOperator::Plus, - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + left_expr: empty, + right_expr: empty, evaluator: Some( binary_create(Cow::Owned(LogicalType::Integer), BinaryOperator::Plus).unwrap(), ), @@ -1363,8 +1211,8 @@ mod test { &mut cursor, ScalarExpression::Binary { op: BinaryOperator::Plus, - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + left_expr: empty, + right_expr: empty, evaluator: None, ty: LogicalType::Integer, }, @@ -1377,7 +1225,7 @@ mod test { ScalarExpression::AggCall { distinct: true, kind: AggKind::Avg, - args: vec![ScalarExpression::Empty], + args: vec![empty], ty: LogicalType::Double, }, Some(&context), @@ -1388,8 +1236,8 @@ mod test { &mut cursor, ScalarExpression::In { negated: true, - expr: Box::new(ScalarExpression::Empty), - args: vec![ScalarExpression::Empty], + expr: empty, + args: vec![empty], }, Some(&context), &mut reference_tables, @@ -1399,9 +1247,9 @@ mod test { &mut cursor, ScalarExpression::Between { negated: true, - expr: Box::new(ScalarExpression::Empty), - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + expr: empty, + left_expr: empty, + right_expr: empty, }, Some(&context), &mut reference_tables, @@ -1410,9 +1258,9 @@ mod test { fn_assert( &mut cursor, ScalarExpression::SubString { - expr: Box::new(ScalarExpression::Empty), - for_expr: Some(Box::new(ScalarExpression::Empty)), - from_expr: Some(Box::new(ScalarExpression::Empty)), + expr: empty, + for_expr: Some(empty), + from_expr: Some(empty), }, Some(&context), &mut reference_tables, @@ -1421,9 +1269,9 @@ mod test { fn_assert( &mut cursor, ScalarExpression::SubString { - expr: Box::new(ScalarExpression::Empty), + expr: empty, for_expr: None, - from_expr: Some(Box::new(ScalarExpression::Empty)), + from_expr: Some(empty), }, Some(&context), &mut reference_tables, @@ -1432,7 +1280,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::SubString { - expr: Box::new(ScalarExpression::Empty), + expr: empty, for_expr: None, from_expr: None, }, @@ -1443,8 +1291,8 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Position { - expr: Box::new(ScalarExpression::Empty), - in_expr: Box::new(ScalarExpression::Empty), + expr: empty, + in_expr: empty, }, Some(&context), &mut reference_tables, @@ -1453,8 +1301,8 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Trim { - expr: Box::new(ScalarExpression::Empty), - trim_what_expr: Some(Box::new(ScalarExpression::Empty)), + expr: empty, + trim_what_expr: Some(empty), trim_where: Some(TrimWhereField::Both), }, Some(&context), @@ -1464,7 +1312,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Trim { - expr: Box::new(ScalarExpression::Empty), + expr: empty, trim_what_expr: None, trim_where: Some(TrimWhereField::Both), }, @@ -1475,7 +1323,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Trim { - expr: Box::new(ScalarExpression::Empty), + expr: empty, trim_what_expr: None, trim_where: None, }, @@ -1492,7 +1340,7 @@ mod test { )?; fn_assert( &mut cursor, - ScalarExpression::Tuple(vec![ScalarExpression::Empty]), + ScalarExpression::Tuple(vec![empty]), Some(&context), &mut reference_tables, &mut plan_arena, @@ -1500,7 +1348,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::ScalaFunction(ScalarFunction { - args: vec![ScalarExpression::Empty], + args: vec![empty], inner: ArcScalarFunctionImpl(CurrentDate::new()), }), Some(&context), @@ -1510,7 +1358,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::TableFunction(TableFunction { - args: vec![ScalarExpression::Empty], + args: vec![empty], catalog: TableFunctionCatalog { schema: Vec::new(), inner: ArcTableFunctionImpl(Numbers::new()), @@ -1523,9 +1371,9 @@ mod test { fn_assert( &mut cursor, ScalarExpression::If { - condition: Box::new(ScalarExpression::Empty), - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + condition: empty, + left_expr: empty, + right_expr: empty, ty: LogicalType::Integer, }, Some(&context), @@ -1535,8 +1383,8 @@ mod test { fn_assert( &mut cursor, ScalarExpression::IfNull { - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + left_expr: empty, + right_expr: empty, ty: LogicalType::Integer, }, Some(&context), @@ -1546,8 +1394,8 @@ mod test { fn_assert( &mut cursor, ScalarExpression::NullIf { - left_expr: Box::new(ScalarExpression::Empty), - right_expr: Box::new(ScalarExpression::Empty), + left_expr: empty, + right_expr: empty, ty: LogicalType::Integer, }, Some(&context), @@ -1557,7 +1405,7 @@ mod test { fn_assert( &mut cursor, ScalarExpression::Coalesce { - exprs: vec![ScalarExpression::Empty], + exprs: vec![empty], ty: LogicalType::Integer, }, Some(&context), @@ -1567,21 +1415,24 @@ mod test { fn_assert( &mut cursor, ScalarExpression::CaseWhen { - operand_expr: Some(Box::new(ScalarExpression::Empty)), - expr_pairs: vec![(ScalarExpression::Empty, ScalarExpression::Empty)], - else_expr: Some(Box::new(ScalarExpression::Empty)), + operand_expr: Some(empty), + expr_pairs: vec![(empty, empty)], + else_expr: Some(empty), ty: LogicalType::Integer, }, Some(&context), &mut reference_tables, &mut plan_arena, )?; + let one = plan_arena.alloc_expression(ScalarExpression::Constant(1.into())); + let two = plan_arena.alloc_expression(ScalarExpression::Constant(2.into())); + let three = plan_arena.alloc_expression(ScalarExpression::Constant(3.into())); fn_assert( &mut cursor, ScalarExpression::CaseWhen { operand_expr: None, - expr_pairs: vec![(ScalarExpression::Empty, ScalarExpression::Empty)], - else_expr: Some(Box::new(ScalarExpression::Empty)), + expr_pairs: vec![(empty, empty)], + else_expr: Some(empty), ty: LogicalType::Integer, }, Some(&context), @@ -1592,7 +1443,7 @@ mod test { &mut cursor, ScalarExpression::CaseWhen { operand_expr: None, - expr_pairs: vec![(ScalarExpression::Empty, ScalarExpression::Empty)], + expr_pairs: vec![(empty, empty)], else_expr: None, ty: LogicalType::Integer, }, @@ -1605,12 +1456,12 @@ mod test { ScalarExpression::WindowCall(WindowCall { function: WindowFunction { kind: WindowFunctionKind::Aggregate(AggKind::Sum), - args: vec![ScalarExpression::Constant(1.into())], + args: vec![one], ty: LogicalType::Integer, }, spec: WindowSpec { - partition_by: vec![ScalarExpression::Constant(2.into())], - order_by: vec![ScalarExpression::Constant(3.into()).desc()], + partition_by: vec![two], + order_by: vec![crate::planner::operator::sort::SortField::from(three).desc()], }, }), Some(&context), diff --git a/src/expression/range_detacher.rs b/src/expression/range_detacher.rs index 8dc26b01..d1ccf29f 100644 --- a/src/expression/range_detacher.rs +++ b/src/expression/range_detacher.rs @@ -16,7 +16,8 @@ use crate::catalog::ColumnRef; use crate::errors::DatabaseError; use crate::expression::{BinaryOperator, ScalarExpression}; use crate::iter_ext::Itertools; -use crate::planner::PlanArena; +use crate::planner::{ExprRef, PlanArena}; +use crate::types::index::IndexMetaRef; use crate::types::value::DataValue; use crate::types::{ColumnId, LogicalType}; use kite_sql_serde_macros::ReferenceSerialization; @@ -42,7 +43,7 @@ pub enum Range { #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct DetachedPredicate { pub(crate) range: Range, - pub(crate) residual: Option, + pub(crate) residual: Option, } impl DetachedPredicate { @@ -54,17 +55,18 @@ impl DetachedPredicate { } fn combine_residuals( - left: Option, - right: Option, - ) -> Option { + left: Option, + right: Option, + arena: &mut PlanArena<'_>, + ) -> Option { match (left, right) { - (Some(left), Some(right)) => Some(ScalarExpression::Binary { + (Some(left), Some(right)) => Some(arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(left), - right_expr: Box::new(right), + left_expr: left, + right_expr: right, evaluator: None, ty: LogicalType::Boolean, - }), + })), (Some(expr), None) | (None, Some(expr)) => Some(expr), (None, None) => None, } @@ -208,34 +210,52 @@ impl Range { for tuple in combinations { collect_tuple_range(&mut ranges, &tuple, self.clone()) } - Some(RangeDetacher::ranges2range(ranges)) + Some(RangeDetacher::::ranges2range(ranges)) } } -pub struct RangeDetacher<'a, 'p> { - table_name: &'a str, - column_id: &'a ColumnId, - arena: &'a PlanArena<'p>, +pub trait RangeColumnMatcher { + fn matches(&self, table_name: &str, column_id: ColumnId, arena: &PlanArena<'_>) -> bool; } -impl<'a, 'p> RangeDetacher<'a, 'p> { - pub(crate) fn new( - table_name: &'a str, - column_id: &'a ColumnId, - arena: &'a PlanArena<'p>, +pub struct IndexRangeColumn { + meta: IndexMetaRef, + position: usize, +} + +impl RangeColumnMatcher for IndexRangeColumn { + fn matches(&self, table_name: &str, column_id: ColumnId, arena: &PlanArena<'_>) -> bool { + let index = arena.index(self.meta); + table_name == index.table_name.as_ref() + && index.column_ids.get(self.position) == Some(&column_id) + } +} + +pub struct RangeDetacher<'a, 'p, M: RangeColumnMatcher = IndexRangeColumn> { + column: M, + arena: &'a mut PlanArena<'p>, +} + +impl<'a, 'p> RangeDetacher<'a, 'p, IndexRangeColumn> { + pub(crate) fn for_index( + meta: IndexMetaRef, + position: usize, + arena: &'a mut PlanArena<'p>, ) -> Self { Self { - table_name, - column_id, + column: IndexRangeColumn { meta, position }, arena, } } +} +impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { pub(crate) fn detach( &mut self, - expr: &ScalarExpression, + expr: ExprRef, ) -> Result, DatabaseError> { - Ok(match expr { + let expression = self.arena.expression(expr).clone(); + Ok(match expression { ScalarExpression::Binary { left_expr, right_expr, @@ -243,18 +263,22 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { .. } => { if let (Some(col), Some(val)) = ( - left_expr.unpack_bound_col(false).map(|(column, _)| column), - right_expr.unpack_val(), + left_expr + .unpack_bound_col(self.arena, false) + .map(|(column, _)| column), + right_expr.unpack_val(self.arena), ) { return self - .new_range(*op, col, val, false) + .new_range(op, col, val, false) .map(|range| range.map(DetachedPredicate::consumed)); } else if let (Some(val), Some(col)) = ( - left_expr.unpack_val(), - right_expr.unpack_bound_col(false).map(|(column, _)| column), + left_expr.unpack_val(self.arena), + right_expr + .unpack_bound_col(self.arena, false) + .map(|(column, _)| column), ) { return self - .new_range(*op, col, val, true) + .new_range(op, col, val, true) .map(|range| range.map(DetachedPredicate::consumed)); } @@ -265,27 +289,30 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { let (range, residual) = match (left, right) { (Some(left_range), Some(right_range)) => { let Some(range) = - Self::merge_binary(*op, left_range.range, right_range.range) + Self::merge_binary(op, left_range.range, right_range.range) else { return Ok(None); }; let residual = DetachedPredicate::combine_residuals( left_range.residual, right_range.residual, + self.arena, ); (range, residual) } (Some(detached), None) => { let residual = DetachedPredicate::combine_residuals( detached.residual, - Some(right_expr.as_ref().clone()), + Some(right_expr), + self.arena, ); (detached.range, residual) } (None, Some(detached)) => { let residual = DetachedPredicate::combine_residuals( - Some(left_expr.as_ref().clone()), + Some(left_expr), detached.residual, + self.arena, ); (detached.range, residual) } @@ -298,8 +325,7 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { let right = self.detach(right_expr)?; if let (Some(left), Some(right)) = (left, right) { if left.residual.is_none() && right.residual.is_none() { - if let Some(range) = - Self::merge_binary(*op, left.range, right.range) + if let Some(range) = Self::merge_binary(op, left.range, right.range) { return Ok(Some(DetachedPredicate::consumed(range))); } @@ -313,20 +339,18 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { ScalarExpression::Alias { expr, .. } | ScalarExpression::TypeCast { expr, .. } => { self.detach(expr)? } - ScalarExpression::IsNull { expr, negated, .. } => match expr.as_ref() { + ScalarExpression::IsNull { expr, negated, .. } => match self.arena.expression(expr) { ScalarExpression::ColumnRef { column, .. } => { - let column = self.arena.column(*column); - if let (Some(col_id), Some(col_table)) = (column.id(), column.table_name()) { - if &col_id == self.column_id && col_table.as_ref() == self.table_name { - return Ok(if *negated { - Some(DetachedPredicate::consumed(Range::Scope { - min: Bound::Unbounded, - max: Bound::Excluded(DataValue::Null), - })) - } else { - Some(DetachedPredicate::consumed(Range::Eq(DataValue::Null))) - }); - } + let column = *column; + if self.matches_column(column) { + return Ok(if negated { + Some(DetachedPredicate::consumed(Range::Scope { + min: Bound::Unbounded, + max: Bound::Excluded(DataValue::Null), + })) + } else { + Some(DetachedPredicate::consumed(Range::Eq(DataValue::Null))) + }); } None @@ -776,13 +800,14 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { } } - fn _is_belong(&self, col: ColumnRef) -> bool { - let col = self.arena.column(col); - matches!( - col.table_name() - .map(|name| self.table_name == name.as_ref()), - Some(true) - ) + fn matches_column(&self, col: ColumnRef) -> bool { + let column = self.arena.column(col); + let (Some(column_id), Some(table_name)) = (column.id(), column.table_name()) else { + return false; + }; + + self.column + .matches(table_name.as_ref(), column_id, self.arena) } fn bound_compared( @@ -825,10 +850,10 @@ impl<'a, 'p> RangeDetacher<'a, 'p> { mut val: DataValue, is_flip: bool, ) -> Result, DatabaseError> { - let column = self.arena.column(col); - if !self._is_belong(col) || column.id() != Some(*self.column_id) { + if !self.matches_column(col) { return Ok(None); } + let column = self.arena.column(col); if val.is_null() { return Ok(match op { BinaryOperator::Spaceship => Some(Range::Eq(DataValue::Null)), @@ -912,20 +937,52 @@ impl fmt::Display for Range { } } +#[cfg(test)] +pub(crate) mod test_support { + use super::*; + + pub(crate) struct DirectRangeColumn<'a> { + table_name: &'a str, + column_id: &'a ColumnId, + } + + impl RangeColumnMatcher for DirectRangeColumn<'_> { + fn matches(&self, table_name: &str, column_id: ColumnId, _arena: &PlanArena<'_>) -> bool { + table_name == self.table_name && column_id == *self.column_id + } + } + + impl<'a, 'p> RangeDetacher<'a, 'p, DirectRangeColumn<'a>> { + pub(crate) fn new( + table_name: &'a str, + column_id: &'a ColumnId, + arena: &'a mut PlanArena<'p>, + ) -> Self { + Self { + column: DirectRangeColumn { + table_name, + column_id, + }, + arena, + } + } + } +} + #[cfg(all(test, not(target_arch = "wasm32")))] #[allow(clippy::uninlined_format_args)] mod test { use crate::binder::test::build_t1_table; use crate::catalog::{ColumnCatalog, ColumnDesc, ColumnRef, TableName}; use crate::errors::DatabaseError; - use crate::expression::range_detacher::{Range, RangeDetacher}; + use crate::expression::range_detacher::{IndexRangeColumn, Range, RangeDetacher}; use crate::expression::{BinaryOperator, ScalarExpression}; use crate::optimizer::heuristic::batch::HepBatchStrategy; use crate::optimizer::heuristic::optimizer::HepOptimizerPipeline; use crate::optimizer::rule::normalization::NormalizationRuleImpl; use crate::planner::operator::filter::FilterOperator; use crate::planner::operator::Operator; - use crate::planner::LogicalPlan; + use crate::planner::{ExprRef, LogicalPlan}; use crate::types::evaluator::binary_create; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -966,35 +1023,51 @@ mod test { Ok(arena.alloc_column(column)) } - fn cmp_predicate(column: ColumnRef, op: BinaryOperator, value: i32) -> ScalarExpression { - ScalarExpression::Binary { + fn cmp_predicate( + arena: &mut crate::planner::PlanArena, + column: ColumnRef, + op: BinaryOperator, + value: i32, + ) -> ExprRef { + let left_expr = arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let right_expr = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(value))); + arena.alloc_expression(ScalarExpression::Binary { op, - left_expr: Box::new(ScalarExpression::column_expr(column, 0)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(value))), + left_expr, + right_expr, evaluator: None, ty: LogicalType::Boolean, - } + }) } - fn and_predicate(left: ScalarExpression, right: ScalarExpression) -> ScalarExpression { - ScalarExpression::Binary { + fn and_predicate( + arena: &mut crate::planner::PlanArena, + left: ExprRef, + right: ExprRef, + ) -> ExprRef { + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(left), - right_expr: Box::new(right), + left_expr: left, + right_expr: right, evaluator: None, ty: LogicalType::Boolean, - } + }) } #[test] fn test_detach_consumes_and_predicates_with_residual() -> Result<(), DatabaseError> { let table_state = build_t1_table()?; let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); - let plan = table_state.plan("select * from t1 where c1 > 10 and c2 > 20")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 > 10 and c2 > 20", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let detached = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .expect("c1 predicate should be consumed"); + let detached = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .expect("c1 predicate should be consumed"); assert_eq!( detached.range, @@ -1005,8 +1078,8 @@ mod test { ); let residual = detached.residual.expect("c2 predicate should remain"); let residual_detached = - RangeDetacher::new("t1", table_state.column_id_by_name("c2"), &plan_arena) - .detach(&residual)? + RangeDetacher::new("t1", table_state.column_id_by_name("c2"), &mut plan_arena) + .detach(residual)? .expect("residual should be exactly the c2 range predicate"); assert_eq!( residual_detached.range, @@ -1027,13 +1100,12 @@ mod test { let table_name: TableName = ::std::sync::Arc::from("nullable_t"); let column_id = 1; let column = test_column(&mut plan_arena, &table_name, column_id, "c1", true)?; - let predicate = and_predicate( - cmp_predicate(column, BinaryOperator::Gt, 0), - cmp_predicate(column, BinaryOperator::Lt, 8), - ); + let left = cmp_predicate(&mut plan_arena, column, BinaryOperator::Gt, 0); + let right = cmp_predicate(&mut plan_arena, column, BinaryOperator::Lt, 8); + let predicate = and_predicate(&mut plan_arena, left, right); - let detached = RangeDetacher::new(table_name.as_ref(), &column_id, &plan_arena) - .detach(&predicate)? + let detached = RangeDetacher::new(table_name.as_ref(), &column_id, &mut plan_arena) + .detach(predicate)? .expect("nullable range predicate should be consumed"); assert_eq!( @@ -1057,37 +1129,41 @@ mod test { let column = test_column(&mut plan_arena, &table_name, column_id, "c1", true)?; let cases = [ ( - cmp_predicate(column, BinaryOperator::Gt, 0), + cmp_predicate(&mut plan_arena, column, BinaryOperator::Gt, 0), Range::Scope { min: Bound::Excluded(DataValue::Int32(0)), max: Bound::Excluded(DataValue::Null), }, ), ( - cmp_predicate(column, BinaryOperator::GtEq, 0), + cmp_predicate(&mut plan_arena, column, BinaryOperator::GtEq, 0), Range::Scope { min: Bound::Included(DataValue::Int32(0)), max: Bound::Excluded(DataValue::Null), }, ), ( - cmp_predicate(column, BinaryOperator::Lt, 8), + cmp_predicate(&mut plan_arena, column, BinaryOperator::Lt, 8), Range::Scope { min: Bound::Unbounded, max: Bound::Excluded(DataValue::Int32(8)), }, ), ( - cmp_predicate(column, BinaryOperator::LtEq, 8), + cmp_predicate(&mut plan_arena, column, BinaryOperator::LtEq, 8), Range::Scope { min: Bound::Unbounded, max: Bound::Included(DataValue::Int32(8)), }, ), ( - ScalarExpression::IsNull { - negated: true, - expr: Box::new(ScalarExpression::column_expr(column, 0)), + { + let expr = + plan_arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + plan_arena.alloc_expression(ScalarExpression::IsNull { + negated: true, + expr, + }) }, Range::Scope { min: Bound::Unbounded, @@ -1097,8 +1173,8 @@ mod test { ]; for (predicate, expected) in cases { - let detached = RangeDetacher::new(table_name.as_ref(), &column_id, &plan_arena) - .detach(&predicate)? + let detached = RangeDetacher::new(table_name.as_ref(), &column_id, &mut plan_arena) + .detach(predicate)? .expect("nullable single-sided predicate should be consumed"); assert_eq!(detached.range, expected); @@ -1112,11 +1188,13 @@ mod test { fn test_detach_consumes_complete_or_only_when_both_sides_match() -> Result<(), DatabaseError> { let table_state = build_t1_table()?; let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); - let plan = table_state.plan("select * from t1 where c1 = 1 or c1 = 2")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 or c1 = 2", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let detached = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .expect("both OR branches should be consumed"); + let detached = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .expect("both OR branches should be consumed"); assert_eq!( detached.range, @@ -1133,10 +1211,12 @@ mod test { fn test_detach_does_not_partially_consume_or() -> Result<(), DatabaseError> { let table_state = build_t1_table()?; let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); - let plan = table_state.plan("select * from t1 where c1 = 1 or c2 = 2")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 or c2 = 2", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let detached = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)?; + let detached = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)?; assert_eq!(detached, None); Ok(()) @@ -1147,41 +1227,49 @@ mod test { let table_state = build_t1_table()?; let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); { - let plan = table_state.plan("select * from t1 where c1 = 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1 => {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 = 1.0")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 = 1.0", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1.0 => {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 != 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 != 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 != 1 => {:#?}", range); assert_eq!(range, None) } { - let plan = table_state.plan("select * from t1 where c1 > 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 > 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 > 1 => c1: {}", range); assert_eq!( range, @@ -1192,12 +1280,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 >= 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 >= 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 >= 1 => c1: {}", range); assert_eq!( range, @@ -1208,12 +1298,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 < 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 < 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 < 1 => c1: {}", range); assert_eq!( range, @@ -1224,12 +1316,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 <= 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 <= 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 <= 1 => c1: {}", range); assert_eq!( range, @@ -1240,12 +1334,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 < 1 and c1 >= 0")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 < 1 and c1 >= 0", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 < 1 and c1 >= 0 => c1: {}", range); assert_eq!( range, @@ -1256,12 +1352,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 < 1 or c1 >= 0")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 < 1 or c1 >= 0", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 < 1 or c1 >= 0 => c1: {}", range); assert_eq!( range, @@ -1273,22 +1371,26 @@ mod test { } // and & or { - let plan = table_state.plan("select * from t1 where c1 = 1 and c1 = 0")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 and c1 = 0", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1 and c1 = 0 => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan("select * from t1 where c1 = 1 or c1 = 0")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 or c1 = 0", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1 or c1 = 0 => c1: {}", range); assert_eq!( range, @@ -1299,53 +1401,63 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 = 1 and c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 and c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1 and c1 = 1 => c1: {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 = 1 or c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 = 1 or c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = 1 or c1 = 1 => c1: {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 > 1 and c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 > 1 and c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 > 1 and c1 = 1 => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan("select * from t1 where c1 >= 1 and c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 >= 1 and c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 >= 1 and c1 = 1 => c1: {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 > 1 or c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 > 1 or c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 > 1 or c1 = 1 => c1: {}", range); assert_eq!( range, @@ -1356,12 +1468,14 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 >= 1 or c1 = 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 >= 1 or c1 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 >= 1 or c1 = 1 => c1: {}", range); assert_eq!( range, @@ -1373,13 +1487,16 @@ mod test { } // scope { - let plan = table_state - .plan("select * from t1 where (c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where (c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "(c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4) => c1: {}", range @@ -1393,13 +1510,16 @@ mod test { ) } { - let plan = table_state - .plan("select * from t1 where (c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where (c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "(c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4) => c1: {}", range @@ -1414,14 +1534,16 @@ mod test { } { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where ((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0", + &mut plan_arena, )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0 => c1: {}", range @@ -1429,14 +1551,16 @@ mod test { assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where ((c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)) and c1 = 0", + &mut plan_arena, )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "((c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)) and c1 = 0 => c1: {}", range @@ -1444,14 +1568,16 @@ mod test { assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where ((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) or c1 = 0", + &mut plan_arena, )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) or c1 = 0 => c1: {}", range @@ -1468,14 +1594,16 @@ mod test { ) } { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where ((c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)) or c1 = 0", + &mut plan_arena, )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "((c1 > 0 and c1 < 3) or (c1 > 1 and c1 < 4)) or c1 = 0 => c1: {}", range @@ -1490,22 +1618,24 @@ mod test { } { - let plan = table_state.plan("select * from t1 where (((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) and (c1 >= 0 and c1 <= 2)")?; + let plan = table_state.plan_with_arena("select * from t1 where (((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) and (c1 >= 0 and c1 <= 2)", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("(((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) and (c1 >= 0 and c1 <= 2) => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan("select * from t1 where (((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) or (c1 >= 0 and c1 <= 2)")?; + let plan = table_state.plan_with_arena("select * from t1 where (((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) or (c1 >= 0 and c1 <= 2)", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("(((c1 > 0 and c1 < 3) and (c1 > 1 and c1 < 4)) and c1 = 0) or (c1 >= 0 and c1 <= 2) => c1: {}", range); assert_eq!( range, @@ -1517,12 +1647,13 @@ mod test { } // ranges and ranges { - let plan = table_state.plan("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -1539,12 +1670,13 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -1562,52 +1694,62 @@ mod test { } // empty { - let plan = table_state.plan("select * from t1 where true")?; + let plan = + table_state.plan_with_arena("select * from t1 where true", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("empty => c1: {:#?}", range); assert_eq!(range, None) } // other column { - let plan = table_state.plan("select * from t1 where c2 = 1")?; + let plan = + table_state.plan_with_arena("select * from t1 where c2 = 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c2 = 1 => c1: {:#?}", range); assert_eq!(range, None) } { - let plan = table_state.plan("select * from t1 where c1 > 1 or c2 > 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 > 1 or c2 > 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 > 1 or c2 > 1 => c1: {:#?}", range); assert_eq!(range, None) } { - let plan = table_state.plan("select * from t1 where c1 > c2 or c2 > 1")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 > c2 or c2 > 1", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 > c2 or c2 > 1 => c1: {:#?}", range); assert_eq!(range, None) } // case 1 { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where c1 = 5 or (c1 > 5 and (c1 > 6 or c1 < 8) and c1 < 12)", + &mut plan_arena, )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!( "c1 = 5 or (c1 > 5 and (c1 > 6 or c1 < 8) and c1 < 12) => c1: {}", range @@ -1622,13 +1764,14 @@ mod test { } // case 2 { - let plan = table_state.plan( + let plan = table_state.plan_with_arena( "select * from t1 where ((c2 >= -8 and -4 >= c1) or (c1 >= 0 and 5 > c2)) and ((c2 > 0 and c1 <= 1) or (c1 > -8 and c2 < -6))", - )?; + &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!( "((c2 >= -8 and -4 >= c1) or (c1 >= 0 and 5 > c2)) and ((c2 > 0 and c1 <= 1) or (c1 > -8 and c2 < -6)) => c1: {:#?}", range @@ -1644,11 +1787,11 @@ mod test { let table_state = build_t1_table()?; let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); let mut detach_c1 = |sql: &str| -> Result, DatabaseError> { - let plan = table_state.plan(sql)?; + let plan = table_state.plan_with_arena(sql, &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); Ok( - RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? .map(|detached| detached.range), ) }; @@ -1692,32 +1835,42 @@ mod test { let mut plan_arena = crate::planner::PlanArena::new(&table_state.table_arena); // eq { - let plan = table_state.plan("select * from t1 where c1 = null")?; + let plan = + table_state.plan_with_arena("select * from t1 where c1 = null", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan("select * from t1 where c1 = null or c1 = 1")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null or c1 = 1", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null or c1 = 1 => c1: {}", range); assert_eq!(range, Range::Eq(DataValue::Int32(1))) } { - let plan = table_state.plan("select * from t1 where c1 = null or c1 < 5")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null or c1 < 5", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null or c1 < 5 => c1: {}", range); assert_eq!( range, @@ -1728,13 +1881,16 @@ mod test { ) } { - let plan = - table_state.plan("select * from t1 where c1 = null or (c1 > 1 and c1 < 5)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null or (c1 > 1 and c1 < 5)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null or (c1 > 1 and c1 < 5) => c1: {}", range); assert_eq!( range, @@ -1745,51 +1901,68 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 = null and c1 < 5")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null and c1 < 5", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null and c1 < 5 => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = - table_state.plan("select * from t1 where c1 = null and (c1 > 1 and c1 < 5)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null and (c1 > 1 and c1 < 5)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 = null and (c1 > 1 and c1 < 5) => c1: {}", range); assert_eq!(range, Range::Dummy) } // noteq { - let plan = table_state.plan("select * from t1 where c1 != null")?; + let plan = table_state + .plan_with_arena("select * from t1 where c1 != null", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 != null => c1: {:#?}", range); assert_eq!(range, Some(Range::Dummy)) } { - let plan = table_state.plan("select * from t1 where c1 = null or c1 != 1")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 = null or c1 != 1", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 = null or c1 != 1 => c1: {:#?}", range); assert_eq!(range, None) } { - let plan = table_state.plan("select * from t1 where c1 != null or c1 < 5")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 != null or c1 < 5", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 != null or c1 < 5 => c1: {:#?}", range); assert_eq!( range, @@ -1800,12 +1973,15 @@ mod test { ) } { - let plan = - table_state.plan("select * from t1 where c1 != null or (c1 > 1 and c1 < 5)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 != null or (c1 > 1 and c1 < 5)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range); println!("c1 != null or (c1 > 1 and c1 < 5) => c1: {:#?}", range); assert_eq!( range, @@ -1816,33 +1992,41 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where c1 != null and c1 < 5")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 != null and c1 < 5", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 != null and c1 < 5 => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = - table_state.plan("select * from t1 where c1 != null and (c1 > 1 and c1 < 5)")?; + let plan = table_state.plan_with_arena( + "select * from t1 where c1 != null and (c1 > 1 and c1 < 5)", + &mut plan_arena, + )?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("c1 != null and (c1 > 1 and c1 < 5) => c1: {}", range); assert_eq!(range, Range::Dummy) } { - let plan = table_state.plan("select * from t1 where (c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where (c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("(c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -1859,12 +2043,13 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) or (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -1881,12 +2066,13 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where (c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where (c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("(c1 = null or (c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and ((c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -1903,12 +2089,13 @@ mod test { ) } { - let plan = table_state.plan("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))")?; + let plan = table_state.plan_with_arena("select * from t1 where ((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5))", &mut plan_arena)?; let op = plan_filter(plan, &mut plan_arena)?.unwrap(); - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &plan_arena) - .detach(&op.predicate)? - .map(|detached| detached.range) - .unwrap(); + let range = + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut plan_arena) + .detach(op.predicate)? + .map(|detached| detached.range) + .unwrap(); println!("((c1 < 2 and c1 > 0) or (c1 < 6 and c1 > 4)) and (c1 = null or (c1 < 3 and c1 > 1) or (c1 < 7 and c1 > 5)) => c1: {}", range); assert_eq!( range, @@ -2364,7 +2551,7 @@ mod test { max: Bound::Unbounded, }; assert_eq!( - RangeDetacher::merge_binary( + RangeDetacher::::merge_binary( BinaryOperator::Or, gt_one.clone(), Range::Eq(DataValue::Int32(1)), @@ -2375,7 +2562,7 @@ mod test { }) ); assert_eq!( - RangeDetacher::merge_binary( + RangeDetacher::::merge_binary( BinaryOperator::And, gt_one, Range::Eq(DataValue::Int32(1)), @@ -2383,7 +2570,7 @@ mod test { Some(Range::Dummy) ); - let disjoint = RangeDetacher::merge_binary( + let disjoint = RangeDetacher::::merge_binary( BinaryOperator::Or, Range::Scope { min: Bound::Included(DataValue::Int32(1)), diff --git a/src/expression/simplify.rs b/src/expression/simplify.rs index e6ffbaba..5379e65d 100644 --- a/src/expression/simplify.rs +++ b/src/expression/simplify.rs @@ -14,14 +14,13 @@ use crate::catalog::ColumnRef; use crate::errors::DatabaseError; -use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut}; -use crate::expression::{BinaryOperator, ScalarExpression, UnaryOperator}; -use crate::planner::PlanArena; +use crate::expression::visitor_mut::ExprVisitorMut; +use crate::expression::{BinaryOperator, ScalarExpression, TypeCast, UnaryOperator}; +use crate::planner::{ExprRef, PlanArena}; use crate::types::evaluator::{binary_create, unary_create}; use crate::types::value::DataValue; use crate::types::LogicalType; use std::borrow::Cow; -use std::mem; #[derive(Debug)] enum Replace { @@ -31,8 +30,8 @@ enum Replace { #[derive(Debug)] struct ReplaceBinary { - column_expr: ScalarExpression, - val_expr: ScalarExpression, + column_expr: ExprRef, + val_expr: ExprRef, op: BinaryOperator, ty: LogicalType, is_column_left: bool, @@ -40,23 +39,25 @@ struct ReplaceBinary { #[derive(Debug)] struct ReplaceUnary { - child_expr: ScalarExpression, + child_expr: ExprRef, op: UnaryOperator, ty: LogicalType, } -pub struct ConstantCalculator<'a, 'p> { - arena: &'a PlanArena<'p>, -} +pub struct ConstantCalculator; -impl<'a, 'p> ConstantCalculator<'a, 'p> { - pub fn new(arena: &'a PlanArena<'p>) -> Self { - Self { arena } +impl ConstantCalculator { + pub fn new(_arena: &PlanArena<'_>) -> Self { + Self } } -impl ExprVisitorMut<'_> for ConstantCalculator<'_, '_> { - fn visit(&mut self, expr: &'_ mut ScalarExpression) -> Result<(), DatabaseError> { +impl ExprVisitorMut for ConstantCalculator { + fn visit_expression( + &mut self, + expr: &mut ScalarExpression, + arena: &mut PlanArena<'_>, + ) -> Result { match expr { ScalarExpression::Unary { op, @@ -64,15 +65,15 @@ impl ExprVisitorMut<'_> for ConstantCalculator<'_, '_> { evaluator, ty, } => { - self.visit(arg_expr)?; + self.visit(arg_expr, arena)?; - if let ScalarExpression::Constant(unary_val) = arg_expr.as_ref() { + if let ScalarExpression::Constant(unary_val) = arena.expression(*arg_expr) { let value = if let Some(evaluator) = evaluator { evaluator.unary_eval(unary_val) } else { unary_create(Cow::Borrowed(ty), *op)?.unary_eval(unary_val) }; - let _ = mem::replace(expr, ScalarExpression::Constant(value)); + *expr = ScalarExpression::Constant(value); } } ScalarExpression::Binary { @@ -81,39 +82,40 @@ impl ExprVisitorMut<'_> for ConstantCalculator<'_, '_> { right_expr, .. } => { - let left_ty = left_expr.return_type(self.arena); - let right_ty = right_expr.return_type(self.arena); - let ty = LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned(); - self.visit(left_expr)?; - self.visit(right_expr)?; + let ty = LogicalType::max_logical_type( + &left_expr.return_type(arena), + &right_expr.return_type(arena), + )? + .into_owned(); + self.visit(left_expr, arena)?; + self.visit(right_expr, arena)?; if let ( ScalarExpression::Constant(left_val), ScalarExpression::Constant(right_val), - ) = (left_expr.as_mut(), right_expr.as_mut()) + ) = (arena.expression(*left_expr), arena.expression(*right_expr)) { let evaluator = binary_create(Cow::Borrowed(&ty), *op)?; - - *left_val = mem::replace(left_val, DataValue::Null).cast(&ty)?; - *right_val = mem::replace(right_val, DataValue::Null).cast(&ty)?; - let value = evaluator.binary_eval(left_val, right_val)?; - let _ = mem::replace(expr, ScalarExpression::Constant(value)); + let left_val = left_val.clone().cast(&ty)?; + let right_val = right_val.clone().cast(&ty)?; + let value = evaluator.binary_eval(&left_val, &right_val)?; + *expr = ScalarExpression::Constant(value); } } ScalarExpression::TypeCast { expr: arg_expr, ty, .. } => { - self.visit(arg_expr)?; + self.visit(arg_expr, arena)?; - if let ScalarExpression::Constant(value) = arg_expr.as_mut() { - let casted = mem::replace(value, DataValue::Null).cast(ty)?; - let _ = mem::replace(expr, ScalarExpression::Constant(casted)); + if let ScalarExpression::Constant(value) = arena.expression(*arg_expr) { + let casted = value.clone().cast(ty)?; + *expr = ScalarExpression::Constant(casted); } } - _ => walk_mut_expr(self, expr)?, + _ => return Ok(true), } - Ok(()) + Ok(false) } } @@ -122,8 +124,12 @@ pub struct Simplify { replaces: Vec, } -impl ExprVisitorMut<'_> for Simplify { - fn visit(&mut self, expr: &'_ mut ScalarExpression) -> Result<(), DatabaseError> { +impl ExprVisitorMut for Simplify { + fn visit_expression( + &mut self, + expr: &mut ScalarExpression, + arena: &mut PlanArena<'_>, + ) -> Result { match expr { ScalarExpression::Unary { op, @@ -133,8 +139,8 @@ impl ExprVisitorMut<'_> for Simplify { } => { let op = *op; let ty = ty.clone(); - let child_expr = arg_expr.as_ref().clone(); - let value = if let Some(value) = arg_expr.unpack_val() { + let child_expr = *arg_expr; + let value = if let Some(value) = arg_expr.unpack_val(arena) { Some(if let Some(evaluator) = evaluator { evaluator.unary_eval(&value) } else { @@ -145,11 +151,11 @@ impl ExprVisitorMut<'_> for Simplify { }; if let Some(value) = value { - let _ = mem::replace(expr, ScalarExpression::Constant(value)); + *expr = ScalarExpression::Constant(value); } else if matches!(op, UnaryOperator::Not) { - if let Some(new_expr) = Self::take_negated_range_comparison(arg_expr) { - let _ = mem::replace(expr, new_expr); - self.visit(expr)?; + if let Some(new_expr) = Self::take_negated_range_comparison(*arg_expr, arena) { + *expr = new_expr; + return self.visit_expression(expr, arena); } else { self.replaces .push(Replace::Unary(ReplaceUnary { child_expr, op, ty })); @@ -166,28 +172,28 @@ impl ExprVisitorMut<'_> for Simplify { ty, .. } => { - self.fix_expr(left_expr, right_expr, op)?; + self.fix_expr(left_expr, right_expr, op, arena)?; // `(c1 - 1) and (c1 + 2)` cannot fix! - self.fix_expr(right_expr, left_expr, op)?; + self.fix_expr(right_expr, left_expr, op, arena)?; if let Some(new_expr) = - Self::take_bool_normalized_range_comparison(*op, left_expr, right_expr) + Self::take_bool_normalized_range_comparison(*op, *left_expr, *right_expr, arena) { - let _ = mem::replace(expr, new_expr); - self.visit(expr)?; - return Ok(()); + *expr = new_expr; + return self.visit_expression(expr, arena); } if Self::is_arithmetic(op) { match ( - left_expr.unpack_bound_col(false), - right_expr.unpack_bound_col(false), + left_expr.unpack_bound_col(arena, false), + right_expr.unpack_bound_col(arena, false), ) { (Some((col, position)), None) => { self.replaces.push(Replace::Binary(ReplaceBinary { - column_expr: ScalarExpression::column_expr(col, position), - val_expr: mem::replace(right_expr, ScalarExpression::Empty), + column_expr: arena + .alloc_expression(ScalarExpression::column_expr(col, position)), + val_expr: *right_expr, op: *op, ty: ty.clone(), is_column_left: true, @@ -195,8 +201,9 @@ impl ExprVisitorMut<'_> for Simplify { } (None, Some((col, position))) => { self.replaces.push(Replace::Binary(ReplaceBinary { - column_expr: ScalarExpression::column_expr(col, position), - val_expr: mem::replace(left_expr, ScalarExpression::Empty), + column_expr: arena + .alloc_expression(ScalarExpression::column_expr(col, position)), + val_expr: *left_expr, op: *op, ty: ty.clone(), is_column_left: false, @@ -204,17 +211,19 @@ impl ExprVisitorMut<'_> for Simplify { } (None, None) => { if self.replaces.is_empty() { - return Ok(()); + return Ok(false); } match ( - left_expr.unpack_bound_col(true), - right_expr.unpack_bound_col(true), + left_expr.unpack_bound_col(arena, true), + right_expr.unpack_bound_col(arena, true), ) { (Some((col, position)), None) => { self.replaces.push(Replace::Binary(ReplaceBinary { - column_expr: ScalarExpression::column_expr(col, position), - val_expr: mem::replace(right_expr, ScalarExpression::Empty), + column_expr: arena.alloc_expression( + ScalarExpression::column_expr(col, position), + ), + val_expr: *right_expr, op: *op, ty: ty.clone(), is_column_left: true, @@ -222,8 +231,10 @@ impl ExprVisitorMut<'_> for Simplify { } (None, Some((col, position))) => { self.replaces.push(Replace::Binary(ReplaceBinary { - column_expr: ScalarExpression::column_expr(col, position), - val_expr: mem::replace(left_expr, ScalarExpression::Empty), + column_expr: arena.alloc_expression( + ScalarExpression::column_expr(col, position), + ), + val_expr: *left_expr, op: *op, ty: ty.clone(), is_column_left: false, @@ -236,17 +247,15 @@ impl ExprVisitorMut<'_> for Simplify { } } } - ScalarExpression::TypeCast { .. } => { - if let Some(val) = expr.unpack_val() { - let _ = mem::replace(expr, ScalarExpression::Constant(val)); + ScalarExpression::TypeCast { expr: arg, ty, .. } => { + if let Some(value) = arg.unpack_val(arena).and_then(|value| value.cast(ty).ok()) { + *expr = ScalarExpression::Constant(value); } } - ScalarExpression::IsNull { .. } => { - if let Some(val) = expr.unpack_val() { - let _ = mem::replace( - expr, - ScalarExpression::Constant(DataValue::Boolean(val.is_null())), - ); + ScalarExpression::IsNull { negated, expr: arg } => { + if let Some(value) = arg.unpack_val(arena) { + *expr = + ScalarExpression::Constant(DataValue::Boolean(value.is_null() != *negated)); } } ScalarExpression::In { @@ -255,7 +264,7 @@ impl ExprVisitorMut<'_> for Simplify { args, } => { if args.is_empty() { - return Ok(()); + return Ok(false); } let (op_1, op_2) = if *negated { @@ -265,8 +274,8 @@ impl ExprVisitorMut<'_> for Simplify { }; let mut new_expr = ScalarExpression::Binary { op: op_1, - left_expr: arg_expr.clone(), - right_expr: Box::new(args.remove(0)), + left_expr: *arg_expr, + right_expr: args.remove(0), evaluator: None, ty: LogicalType::Boolean, }; @@ -274,21 +283,20 @@ impl ExprVisitorMut<'_> for Simplify { for arg in args.drain(..) { new_expr = ScalarExpression::Binary { op: op_2, - left_expr: Box::new(ScalarExpression::Binary { + left_expr: arena.alloc_expression(ScalarExpression::Binary { op: op_1, - left_expr: arg_expr.clone(), - right_expr: Box::new(arg), + left_expr: *arg_expr, + right_expr: arg, evaluator: None, ty: LogicalType::Boolean, }), - right_expr: Box::new(new_expr), + right_expr: arena.alloc_expression(new_expr), evaluator: None, ty: LogicalType::Boolean, - } + }; } - let _ = mem::replace(expr, new_expr); - - walk_mut_expr(self, expr)?; + *expr = new_expr; + return Ok(true); } ScalarExpression::Between { negated, @@ -305,34 +313,30 @@ impl ExprVisitorMut<'_> for Simplify { BinaryOperator::LtEq, ) }; - let new_expr = ScalarExpression::Binary { + *expr = ScalarExpression::Binary { op, - left_expr: Box::new(ScalarExpression::Binary { + left_expr: arena.alloc_expression(ScalarExpression::Binary { op: left_op, - left_expr: arg_expr.clone(), - right_expr: mem::replace(left_expr, Box::new(ScalarExpression::Empty)), + left_expr: *arg_expr, + right_expr: *left_expr, evaluator: None, ty: LogicalType::Boolean, }), - right_expr: Box::new(ScalarExpression::Binary { + right_expr: arena.alloc_expression(ScalarExpression::Binary { op: right_op, - left_expr: mem::replace(arg_expr, Box::new(ScalarExpression::Empty)), - right_expr: mem::replace(right_expr, Box::new(ScalarExpression::Empty)), + left_expr: *arg_expr, + right_expr: *right_expr, evaluator: None, ty: LogicalType::Boolean, }), evaluator: None, ty: LogicalType::Boolean, }; - - let _ = mem::replace(expr, new_expr); - - walk_mut_expr(self, expr)?; + return Ok(true); } - _ => walk_mut_expr(self, expr)?, + _ => return Ok(true), } - - Ok(()) + Ok(false) } } @@ -357,64 +361,73 @@ impl Simplify { } } - fn take_range_comparison(expr: &mut Box) -> Option { - match expr.as_ref() { - ScalarExpression::Binary { op, .. } if Self::negate_range_comparison(*op).is_some() => { - Some(mem::replace(expr.as_mut(), ScalarExpression::Empty)) + fn take_range_comparison(expr: ExprRef, arena: &PlanArena<'_>) -> Option { + match arena.expression(expr) { + expression @ ScalarExpression::Binary { op, .. } + if Self::negate_range_comparison(*op).is_some() => + { + Some(expression.clone()) } _ => None, } } - fn take_negated_range_comparison(expr: &mut Box) -> Option { - match expr.as_mut() { + fn take_negated_range_comparison( + expr: ExprRef, + arena: &PlanArena<'_>, + ) -> Option { + let mut expression = arena.expression(expr).clone(); + match &mut expression { ScalarExpression::Binary { op, .. } => { *op = Self::negate_range_comparison(*op)?; - Some(mem::replace(expr.as_mut(), ScalarExpression::Empty)) + Some(expression) } _ => None, } } - fn boolean_constant(expr: &ScalarExpression) -> Option { - match expr { + fn boolean_constant(expr: ExprRef, arena: &PlanArena<'_>) -> Option { + match arena.expression(expr) { ScalarExpression::Constant(DataValue::Boolean(value)) => Some(*value), _ => None, } } fn take_range_comparison_with_polarity( - expr: &mut Box, + expr: ExprRef, positive: bool, + arena: &PlanArena<'_>, ) -> Option { if positive { - Self::take_range_comparison(expr) + Self::take_range_comparison(expr, arena) } else { - Self::take_negated_range_comparison(expr) + Self::take_negated_range_comparison(expr, arena) } } fn take_bool_normalized_range_comparison( op: BinaryOperator, - left_expr: &mut Box, - right_expr: &mut Box, + left_expr: ExprRef, + right_expr: ExprRef, + arena: &PlanArena<'_>, ) -> Option { let is_eq = matches!(op, BinaryOperator::Eq); - let is_not_eq = matches!(op, BinaryOperator::NotEq); - if !is_eq && !is_not_eq { + if !matches!(op, BinaryOperator::Eq | BinaryOperator::NotEq) { return None; } - if let Some(value) = Self::boolean_constant(right_expr) { + if let Some(value) = Self::boolean_constant(right_expr, arena) { return Self::take_range_comparison_with_polarity( left_expr, if is_eq { value } else { !value }, + arena, ); } - if let Some(value) = Self::boolean_constant(left_expr) { + if let Some(value) = Self::boolean_constant(left_expr, arena) { return Self::take_range_comparison_with_polarity( right_expr, if is_eq { value } else { !value }, + arena, ); } @@ -423,21 +436,24 @@ impl Simplify { fn fix_expr( &mut self, - left_expr: &mut Box, - right_expr: &mut Box, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, op: &mut BinaryOperator, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; + self.visit(left_expr, arena)?; if Self::is_arithmetic(op) { return Ok(()); } while let Some(replace) = self.replaces.pop() { match replace { - Replace::Binary(binary) => Self::fix_binary(binary, left_expr, right_expr, op), + Replace::Binary(binary) => { + Self::fix_binary(binary, left_expr, right_expr, op, arena) + } Replace::Unary(unary) => { - Self::fix_unary(unary, left_expr, right_expr, op); - self.fix_expr(left_expr, right_expr, op)?; + Self::fix_unary(unary, left_expr, right_expr, op, arena); + self.fix_expr(left_expr, right_expr, op, arena)?; } } } @@ -447,58 +463,53 @@ impl Simplify { fn fix_unary( replace_unary: ReplaceUnary, - col_expr: &mut Box, - val_expr: &mut Box, + col_expr: &mut ExprRef, + val_expr: &mut ExprRef, op: &mut BinaryOperator, + arena: &mut PlanArena<'_>, ) { let ReplaceUnary { child_expr, op: fix_op, ty: fix_ty, } = replace_unary; - let _ = mem::replace(col_expr, Box::new(child_expr)); + *col_expr = child_expr; - let expr = mem::replace(val_expr, Box::new(ScalarExpression::Empty)); - let _ = mem::replace( - val_expr, - Box::new(ScalarExpression::Unary { - op: fix_op, - expr, - evaluator: None, - ty: fix_ty, - }), - ); - let _ = mem::replace( - op, - match fix_op { - UnaryOperator::Plus => *op, - UnaryOperator::Minus => match *op { - BinaryOperator::Plus => BinaryOperator::Minus, - BinaryOperator::Minus => BinaryOperator::Plus, - BinaryOperator::Multiply => BinaryOperator::Divide, - BinaryOperator::Divide => BinaryOperator::Multiply, - BinaryOperator::Gt => BinaryOperator::Lt, - BinaryOperator::Lt => BinaryOperator::Gt, - BinaryOperator::GtEq => BinaryOperator::LtEq, - BinaryOperator::LtEq => BinaryOperator::GtEq, - source_op => source_op, - }, - UnaryOperator::Not => match *op { - BinaryOperator::Gt => BinaryOperator::Lt, - BinaryOperator::Lt => BinaryOperator::Gt, - BinaryOperator::GtEq => BinaryOperator::LtEq, - BinaryOperator::LtEq => BinaryOperator::GtEq, - source_op => source_op, - }, + *val_expr = arena.alloc_expression(ScalarExpression::Unary { + op: fix_op, + expr: *val_expr, + evaluator: None, + ty: fix_ty, + }); + *op = match fix_op { + UnaryOperator::Plus => *op, + UnaryOperator::Minus => match *op { + BinaryOperator::Plus => BinaryOperator::Minus, + BinaryOperator::Minus => BinaryOperator::Plus, + BinaryOperator::Multiply => BinaryOperator::Divide, + BinaryOperator::Divide => BinaryOperator::Multiply, + BinaryOperator::Gt => BinaryOperator::Lt, + BinaryOperator::Lt => BinaryOperator::Gt, + BinaryOperator::GtEq => BinaryOperator::LtEq, + BinaryOperator::LtEq => BinaryOperator::GtEq, + source_op => source_op, }, - ); + UnaryOperator::Not => match *op { + BinaryOperator::Gt => BinaryOperator::Lt, + BinaryOperator::Lt => BinaryOperator::Gt, + BinaryOperator::GtEq => BinaryOperator::LtEq, + BinaryOperator::LtEq => BinaryOperator::GtEq, + source_op => source_op, + }, + }; } fn fix_binary( replace_binary: ReplaceBinary, - left_expr: &mut Box, - right_expr: &mut Box, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, op: &mut BinaryOperator, + arena: &mut PlanArena<'_>, ) { let ReplaceBinary { column_expr, @@ -521,70 +532,58 @@ impl Simplify { BinaryOperator::LtEq => BinaryOperator::GtEq, source_op => source_op, }; - let temp_expr = mem::replace(right_expr, Box::new(ScalarExpression::Empty)); let (fixed_op, fixed_left_expr, fixed_right_expr) = if is_column_left { - (op_flip(fix_op), temp_expr, Box::new(val_expr)) + (op_flip(fix_op), *right_expr, val_expr) } else { if matches!(fix_op, BinaryOperator::Minus | BinaryOperator::Multiply) { - let _ = mem::replace(op, comparison_flip(*op)); + *op = comparison_flip(*op); } - (fix_op, Box::new(val_expr), temp_expr) + (fix_op, val_expr, *right_expr) }; - let _ = mem::replace(left_expr, Box::new(column_expr)); - let _ = mem::replace( - right_expr, - Box::new(ScalarExpression::Binary { - op: fixed_op, - left_expr: fixed_left_expr, - right_expr: fixed_right_expr, - evaluator: None, - ty: fix_ty, - }), - ); + *left_expr = column_expr; + *right_expr = arena.alloc_expression(ScalarExpression::Binary { + op: fixed_op, + left_expr: fixed_left_expr, + right_expr: fixed_right_expr, + evaluator: None, + ty: fix_ty, + }); } } -impl ScalarExpression { - pub(crate) fn unpack_val(&self) -> Option { - match self { +impl ExprRef { + pub(crate) fn unpack_val(self, arena: &PlanArena<'_>) -> Option { + match arena.expression(self) { ScalarExpression::Constant(val) => Some(val.clone()), - ScalarExpression::Alias { expr, .. } => expr.unpack_val(), + ScalarExpression::Alias { expr, .. } => expr.unpack_val(arena), ScalarExpression::TypeCast { expr, ty, .. } => { - expr.unpack_val().and_then(|val| val.cast(ty).ok()) + expr.unpack_val(arena).and_then(|val| val.cast(ty).ok()) } - ScalarExpression::IsNull { expr, .. } => expr - .unpack_val() - .map(|val| DataValue::Boolean(val.is_null())), + ScalarExpression::IsNull { negated, expr } => Some(DataValue::Boolean( + expr.unpack_val(arena)?.is_null() != *negated, + )), ScalarExpression::Unary { expr, op, evaluator, ty, - .. - } => { - let value = expr.unpack_val()?; - let unary_value = if let Some(evaluator) = evaluator { - evaluator.unary_eval(&value) - } else { - unary_create(Cow::Borrowed(ty), *op) - .ok()? - .unary_eval(&value) - }; - Some(unary_value) - } + } => Some(if let Some(evaluator) = evaluator { + evaluator.unary_eval(&expr.unpack_val(arena)?) + } else { + unary_create(Cow::Borrowed(ty), *op) + .ok()? + .unary_eval(&expr.unpack_val(arena)?) + }), ScalarExpression::Binary { left_expr, right_expr, op, ty, evaluator, - .. } => { - let mut left = left_expr.unpack_val()?; - let mut right = right_expr.unpack_val()?; - left = left.cast(ty).ok()?; - right = right.cast(ty).ok()?; + let left = left_expr.unpack_val(arena)?.cast(ty).ok()?; + let right = right_expr.unpack_val(arena)?.cast(ty).ok()?; if let Some(evaluator) = evaluator { evaluator.binary_eval(&left, &right) } else { @@ -598,11 +597,15 @@ impl ScalarExpression { } } - pub(crate) fn unpack_bound_col(&self, is_deep: bool) -> Option<(ColumnRef, usize)> { - match self { + pub(crate) fn unpack_bound_col( + self, + arena: &PlanArena<'_>, + is_deep: bool, + ) -> Option<(ColumnRef, usize)> { + match arena.expression(self) { ScalarExpression::ColumnRef { column, position } => Some((*column, *position)), - ScalarExpression::Alias { expr, .. } => expr.unpack_bound_col(is_deep), - ScalarExpression::Unary { expr, .. } => expr.unpack_bound_col(is_deep), + ScalarExpression::Alias { expr, .. } => expr.unpack_bound_col(arena, is_deep), + ScalarExpression::Unary { expr, .. } => expr.unpack_bound_col(arena, is_deep), ScalarExpression::Binary { left_expr, right_expr, @@ -613,8 +616,8 @@ impl ScalarExpression { } left_expr - .unpack_bound_col(true) - .or_else(|| right_expr.unpack_bound_col(true)) + .unpack_bound_col(arena, true) + .or_else(|| right_expr.unpack_bound_col(arena, true)) } _ => None, } diff --git a/src/expression/visitor.rs b/src/expression/visitor.rs index ecd85a3f..e1629623 100644 --- a/src/expression/visitor.rs +++ b/src/expression/visitor.rs @@ -4,7 +4,7 @@ // you may not use this file except in compliance with the License. // You may obtain a copy of the License at // -// http://www.apache.org/licenses/LICENSE-2.0 +// http://www.apache.org/licenses/LICENSE-2.0 // // Unless required by applicable law or agreed to in writing, software // distributed under the License is distributed on an "AS IS" BASIS, @@ -18,89 +18,107 @@ use crate::expression::agg::AggKind; use crate::expression::function::scala::ScalarFunction; use crate::expression::function::table::TableFunction; use crate::expression::window::WindowCall; -use crate::expression::TrimWhereField; -use crate::expression::{AliasType, BinaryOperator, ScalarExpression, UnaryOperator}; +use crate::expression::{ + AliasType, BinaryOperator, ScalarExpression, TrimWhereField, UnaryOperator, +}; +use crate::planner::{ExprRef, MetaArena}; use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef}; use crate::types::value::DataValue; use crate::types::LogicalType; -pub trait ExprVisitor<'a>: Sized { - fn visit(&mut self, expr: &'a ScalarExpression) -> Result<(), DatabaseError> { - walk_expr(self, expr) +pub trait ExprVisitor: Sized { + fn visit(&mut self, expr: ExprRef, arena: &A) -> Result<(), DatabaseError> { + if !self.visit_expression_ref(expr, arena)? { + return Ok(()); + } + if self.visit_expression(arena.expression(expr), arena)? { + walk_expr(self, expr, arena)?; + } + Ok(()) } - fn visit_constant(&mut self, _value: &'a DataValue) -> Result<(), DatabaseError> { - Ok(()) + fn visit_expression_ref(&mut self, _expr: ExprRef, _arena: &A) -> Result { + Ok(true) } - fn visit_column_ref(&mut self, _column: &'a ColumnRef) -> Result<(), DatabaseError> { - Ok(()) + fn visit_expression( + &mut self, + _expr: &ScalarExpression, + _arena: &A, + ) -> Result { + Ok(true) } + fn visit_constant(&mut self, _value: &DataValue) -> Result<(), DatabaseError> { + Ok(()) + } + fn visit_column_ref(&mut self, _column: &ColumnRef) -> Result<(), DatabaseError> { + Ok(()) + } fn visit_alias( &mut self, - expr: &'a ScalarExpression, - ty: &'a AliasType, + expr: ExprRef, + alias: &AliasType, + arena: &A, ) -> Result<(), DatabaseError> { - if let AliasType::Expr(alias_expr) = ty { - self.visit(alias_expr)?; + if let AliasType::Expr(alias_expr) = alias { + self.visit(*alias_expr, arena)?; } - self.visit(expr) + self.visit(expr, arena) } - fn visit_type_cast( &mut self, - expr: &'a ScalarExpression, - _ty: &'a LogicalType, - _evaluator: Option<&'a CastEvaluatorRef>, + expr: ExprRef, + _ty: &LogicalType, + _evaluator: Option<&CastEvaluatorRef>, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } - fn visit_is_null( &mut self, _negated: bool, - expr: &'a ScalarExpression, + expr: ExprRef, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } - fn visit_unary( &mut self, - _op: &'a UnaryOperator, - expr: &'a ScalarExpression, - _evaluator: Option<&'a UnaryEvaluatorRef>, - _ty: &'a LogicalType, + _op: &UnaryOperator, + expr: ExprRef, + _evaluator: Option<&UnaryEvaluatorRef>, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } - fn visit_binary( &mut self, - _op: &'a BinaryOperator, - left_expr: &'a ScalarExpression, - right_expr: &'a ScalarExpression, - _evaluator: Option<&'a BinaryEvaluatorRef>, - _ty: &'a LogicalType, + _op: &BinaryOperator, + left: ExprRef, + right: ExprRef, + _evaluator: Option<&BinaryEvaluatorRef>, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left, arena)?; + self.visit(right, arena) } - fn visit_agg( &mut self, _distinct: bool, - _kind: &'a AggKind, - args: &'a [ScalarExpression], - _ty: &'a LogicalType, + _kind: &AggKind, + args: &[ExprRef], + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { for arg in args { - self.visit(arg)?; + self.visit(*arg, arena)?; } Ok(()) } - - fn visit_window(&mut self, window: &'a WindowCall) -> Result<(), DatabaseError> { + fn visit_window(&mut self, window: &WindowCall, arena: &A) -> Result<(), DatabaseError> { for expr in window .function .args @@ -108,268 +126,252 @@ pub trait ExprVisitor<'a>: Sized { .chain(&window.spec.partition_by) .chain(window.spec.order_by.iter().map(|field| &field.expr)) { - self.visit(expr)?; + self.visit(*expr, arena)?; } Ok(()) } - fn visit_in( &mut self, _negated: bool, - expr: &'a ScalarExpression, - args: &'a [ScalarExpression], + expr: ExprRef, + args: &[ExprRef], + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr)?; + self.visit(expr, arena)?; for arg in args { - self.visit(arg)?; + self.visit(*arg, arena)?; } Ok(()) } - fn visit_between( &mut self, _negated: bool, - expr: &'a ScalarExpression, - left_expr: &'a ScalarExpression, - right_expr: &'a ScalarExpression, + expr: ExprRef, + left: ExprRef, + right: ExprRef, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(expr, arena)?; + self.visit(left, arena)?; + self.visit(right, arena) } - fn visit_substring( &mut self, - expr: &'a ScalarExpression, - for_expr: Option<&'a ScalarExpression>, - from_expr: Option<&'a ScalarExpression>, + expr: ExprRef, + for_expr: Option, + from_expr: Option, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - if let Some(for_expr) = for_expr { - self.visit(for_expr)?; + self.visit(expr, arena)?; + if let Some(expr) = for_expr { + self.visit(expr, arena)?; } - if let Some(from_expr) = from_expr { - self.visit(from_expr)?; + if let Some(expr) = from_expr { + self.visit(expr, arena)?; } Ok(()) } - fn visit_position( &mut self, - expr: &'a ScalarExpression, - in_expr: &'a ScalarExpression, + expr: ExprRef, + in_expr: ExprRef, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - self.visit(in_expr) + self.visit(expr, arena)?; + self.visit(in_expr, arena) } - fn visit_trim( &mut self, - expr: &'a ScalarExpression, - trim_what_expr: Option<&'a ScalarExpression>, - _trim_where: Option<&'a TrimWhereField>, + expr: ExprRef, + trim_what: Option, + _trim_where: Option<&TrimWhereField>, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - if let Some(trim_what_expr) = trim_what_expr { - self.visit(trim_what_expr)?; + self.visit(expr, arena)?; + if let Some(expr) = trim_what { + self.visit(expr, arena)?; } Ok(()) } - fn visit_empty(&mut self) -> Result<(), DatabaseError> { Ok(()) } - - fn visit_reference( - &mut self, - expr: &'a ScalarExpression, - _pos: usize, - ) -> Result<(), DatabaseError> { - self.visit(expr) - } - - fn visit_tuple(&mut self, exprs: &'a [ScalarExpression]) -> Result<(), DatabaseError> { + fn visit_tuple(&mut self, exprs: &[ExprRef], arena: &A) -> Result<(), DatabaseError> { for expr in exprs { - self.visit(expr)?; + self.visit(*expr, arena)?; } Ok(()) } - fn visit_scala_function( &mut self, - scalar_function: &'a ScalarFunction, + function: &ScalarFunction, + arena: &A, ) -> Result<(), DatabaseError> { - for arg in &scalar_function.args { - self.visit(arg)?; + for arg in &function.args { + self.visit(*arg, arena)?; } Ok(()) } - fn visit_table_function( &mut self, - table_function: &'a TableFunction, + function: &TableFunction, + arena: &A, ) -> Result<(), DatabaseError> { - for arg in &table_function.args { - self.visit(arg)?; + for arg in &function.args { + self.visit(*arg, arena)?; } Ok(()) } - fn visit_if( &mut self, - condition: &'a ScalarExpression, - left_expr: &'a ScalarExpression, - right_expr: &'a ScalarExpression, - _ty: &'a LogicalType, + condition: ExprRef, + left: ExprRef, + right: ExprRef, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(condition)?; - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(condition, arena)?; + self.visit(left, arena)?; + self.visit(right, arena) } - fn visit_if_null( &mut self, - left_expr: &'a ScalarExpression, - right_expr: &'a ScalarExpression, - _ty: &'a LogicalType, + left: ExprRef, + right: ExprRef, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left, arena)?; + self.visit(right, arena) } - fn visit_null_if( &mut self, - left_expr: &'a ScalarExpression, - right_expr: &'a ScalarExpression, - _ty: &'a LogicalType, + left: ExprRef, + right: ExprRef, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left, arena)?; + self.visit(right, arena) } - fn visit_coalesce( &mut self, - exprs: &'a [ScalarExpression], - _ty: &'a LogicalType, + exprs: &[ExprRef], + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { for expr in exprs { - self.visit(expr)?; + self.visit(*expr, arena)?; } Ok(()) } - fn visit_case_when( &mut self, - operand_expr: Option<&'a ScalarExpression>, - expr_pairs: &'a [(ScalarExpression, ScalarExpression)], - else_expr: Option<&'a ScalarExpression>, - _ty: &'a LogicalType, + operand: Option, + pairs: &[(ExprRef, ExprRef)], + else_expr: Option, + _ty: &LogicalType, + arena: &A, ) -> Result<(), DatabaseError> { - if let Some(operand_expr) = operand_expr { - self.visit(operand_expr)?; + if let Some(expr) = operand { + self.visit(expr, arena)?; } - for (left_expr, right_expr) in expr_pairs { - self.visit(left_expr)?; - self.visit(right_expr)?; + for (left, right) in pairs { + self.visit(*left, arena)?; + self.visit(*right, arena)?; } - if let Some(else_expr) = else_expr { - self.visit(else_expr)?; + if let Some(expr) = else_expr { + self.visit(expr, arena)?; } Ok(()) } } -pub fn walk_expr<'a, V: ExprVisitor<'a>>( +pub fn walk_expr>( visitor: &mut V, - expr: &'a ScalarExpression, + expr: ExprRef, + arena: &A, ) -> Result<(), DatabaseError> { - match expr { + match arena.expression(expr) { ScalarExpression::Constant(value) => visitor.visit_constant(value), ScalarExpression::ColumnRef { column, .. } => visitor.visit_column_ref(column), - ScalarExpression::Alias { expr, alias } => visitor.visit_alias(expr, alias), + ScalarExpression::Alias { expr, alias } => visitor.visit_alias(*expr, alias, arena), ScalarExpression::TypeCast { expr, ty, evaluator, - } => visitor.visit_type_cast(expr, ty, evaluator.as_ref()), - ScalarExpression::IsNull { negated, expr } => visitor.visit_is_null(*negated, expr), + } => visitor.visit_type_cast(*expr, ty, evaluator.as_ref(), arena), + ScalarExpression::IsNull { negated, expr } => visitor.visit_is_null(*negated, *expr, arena), ScalarExpression::Unary { op, expr, evaluator, ty, - } => visitor.visit_unary(op, expr, evaluator.as_ref(), ty), + } => visitor.visit_unary(op, *expr, evaluator.as_ref(), ty, arena), ScalarExpression::Binary { op, left_expr, right_expr, evaluator, ty, - } => visitor.visit_binary(op, left_expr, right_expr, evaluator.as_ref(), ty), + } => visitor.visit_binary(op, *left_expr, *right_expr, evaluator.as_ref(), ty, arena), ScalarExpression::AggCall { distinct, kind, args, ty, - } => visitor.visit_agg(*distinct, kind, args, ty), + } => visitor.visit_agg(*distinct, kind, args, ty, arena), ScalarExpression::In { negated, expr, args, - } => visitor.visit_in(*negated, expr, args), + } => visitor.visit_in(*negated, *expr, args, arena), ScalarExpression::Between { negated, expr, left_expr, right_expr, - } => visitor.visit_between(*negated, expr, left_expr, right_expr), + } => visitor.visit_between(*negated, *expr, *left_expr, *right_expr, arena), ScalarExpression::SubString { expr, for_expr, from_expr, - } => visitor.visit_substring(expr, for_expr.as_deref(), from_expr.as_deref()), - ScalarExpression::Position { expr, in_expr } => visitor.visit_position(expr, in_expr), + } => visitor.visit_substring(*expr, *for_expr, *from_expr, arena), + ScalarExpression::Position { expr, in_expr } => { + visitor.visit_position(*expr, *in_expr, arena) + } ScalarExpression::Trim { expr, trim_what_expr, trim_where, - } => visitor.visit_trim(expr, trim_what_expr.as_deref(), trim_where.as_ref()), + } => visitor.visit_trim(*expr, *trim_what_expr, trim_where.as_ref(), arena), ScalarExpression::Empty => visitor.visit_empty(), - ScalarExpression::Tuple(exprs) => visitor.visit_tuple(exprs), - ScalarExpression::ScalaFunction(scalar_function) => { - visitor.visit_scala_function(scalar_function) - } - ScalarExpression::TableFunction(table_function) => { - visitor.visit_table_function(table_function) - } + ScalarExpression::Tuple(exprs) => visitor.visit_tuple(exprs, arena), + ScalarExpression::ScalaFunction(function) => visitor.visit_scala_function(function, arena), + ScalarExpression::TableFunction(function) => visitor.visit_table_function(function, arena), ScalarExpression::If { condition, left_expr, right_expr, ty, - } => visitor.visit_if(condition, left_expr, right_expr, ty), + } => visitor.visit_if(*condition, *left_expr, *right_expr, ty, arena), ScalarExpression::IfNull { left_expr, right_expr, ty, - } => visitor.visit_if_null(left_expr, right_expr, ty), + } => visitor.visit_if_null(*left_expr, *right_expr, ty, arena), ScalarExpression::NullIf { left_expr, right_expr, ty, - } => visitor.visit_null_if(left_expr, right_expr, ty), - ScalarExpression::Coalesce { exprs, ty } => visitor.visit_coalesce(exprs, ty), + } => visitor.visit_null_if(*left_expr, *right_expr, ty, arena), + ScalarExpression::Coalesce { exprs, ty } => visitor.visit_coalesce(exprs, ty, arena), ScalarExpression::CaseWhen { operand_expr, expr_pairs, else_expr, ty, - } => visitor.visit_case_when( - operand_expr.as_deref(), - expr_pairs, - else_expr.as_deref(), - ty, - ), - ScalarExpression::WindowCall(window) => visitor.visit_window(window), + } => visitor.visit_case_when(*operand_expr, expr_pairs, *else_expr, ty, arena), + ScalarExpression::WindowCall(window) => visitor.visit_window(window, arena), } } diff --git a/src/expression/visitor_mut.rs b/src/expression/visitor_mut.rs index 1ea298ce..d99978eb 100644 --- a/src/expression/visitor_mut.rs +++ b/src/expression/visitor_mut.rs @@ -18,8 +18,10 @@ use crate::expression::agg::AggKind; use crate::expression::function::scala::ScalarFunction; use crate::expression::function::table::TableFunction; use crate::expression::window::WindowCall; -use crate::expression::TrimWhereField; -use crate::expression::{AliasType, BinaryOperator, ScalarExpression, UnaryOperator}; +use crate::expression::{ + AliasType, BinaryOperator, ScalarExpression, TrimWhereField, UnaryOperator, +}; +use crate::planner::{ExprRef, PlanArena}; use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef}; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -28,11 +30,12 @@ pub(crate) struct PositionShift { pub(crate) delta: isize, } -impl ExprVisitorMut<'_> for PositionShift { +impl ExprVisitorMut for PositionShift { fn visit_column_ref( &mut self, _column: &mut ColumnRef, position: &mut usize, + _arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if self.delta.is_negative() { *position = position.saturating_sub(self.delta.unsigned_abs()); @@ -43,87 +46,133 @@ impl ExprVisitorMut<'_> for PositionShift { } } -pub trait ExprVisitorMut<'a>: Sized { - fn visit(&mut self, expr: &'a mut ScalarExpression) -> Result<(), DatabaseError> { - walk_mut_expr(self, expr) +pub trait ExprVisitorMut: Sized { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + if !self.visit_expression_ref(expr, arena)? { + return Ok(()); + } + + let mut expression = + std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + let result = self.visit_expression(&mut expression, arena); + *arena.expression_mut(*expr) = expression; + if result? { + walk_mut_expr(self, expr, arena)?; + } + Ok(()) + } + + fn visit_expression_ref( + &mut self, + _expr: &mut ExprRef, + _arena: &mut PlanArena<'_>, + ) -> Result { + Ok(true) + } + + fn visit_expression( + &mut self, + _expr: &mut ScalarExpression, + _arena: &mut PlanArena<'_>, + ) -> Result { + Ok(true) } - fn visit_constant(&mut self, _value: &'a mut DataValue) -> Result<(), DatabaseError> { + fn visit_constant( + &mut self, + _value: &mut DataValue, + _arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { Ok(()) } fn visit_column_ref( &mut self, - _column: &'a mut ColumnRef, - _position: &'a mut usize, + _column: &mut ColumnRef, + _position: &mut usize, + _arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { Ok(()) } fn visit_alias( &mut self, - expr: &'a mut ScalarExpression, - ty: &'a mut AliasType, + expr: &mut ExprRef, + alias: &mut AliasType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - if let AliasType::Expr(alias_expr) = ty { - self.visit(alias_expr)?; + if let AliasType::Expr(alias_expr) = alias { + self.visit(alias_expr, arena)?; } - self.visit(expr) + self.visit(expr, arena) } fn visit_type_cast( &mut self, - expr: &'a mut ScalarExpression, - _ty: &'a mut LogicalType, - _evaluator: &'a mut Option, + expr: &mut ExprRef, + _ty: &mut LogicalType, + _evaluator: &mut Option, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } fn visit_is_null( &mut self, _negated: bool, - expr: &'a mut ScalarExpression, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } fn visit_unary( &mut self, - _op: &'a mut UnaryOperator, - expr: &'a mut ScalarExpression, - _evaluator: &'a mut Option, - _ty: &'a mut LogicalType, + _op: &mut UnaryOperator, + expr: &mut ExprRef, + _evaluator: &mut Option, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } fn visit_binary( &mut self, - _op: &'a mut BinaryOperator, - left_expr: &'a mut ScalarExpression, - right_expr: &'a mut ScalarExpression, - _evaluator: &'a mut Option, - _ty: &'a mut LogicalType, + _op: &mut BinaryOperator, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + _evaluator: &mut Option, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left_expr, arena)?; + self.visit(right_expr, arena) } fn visit_agg( &mut self, _distinct: bool, - _kind: &'a mut AggKind, - args: &'a mut [ScalarExpression], - _ty: &'a mut LogicalType, + _kind: &mut AggKind, + args: &mut [ExprRef], + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { for arg in args { - self.visit(arg)?; + self.visit(arg, arena)?; } Ok(()) } - fn visit_window(&mut self, window: &'a mut WindowCall) -> Result<(), DatabaseError> { + fn visit_window( + &mut self, + window: &mut WindowCall, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { for expr in window .function .args @@ -131,7 +180,7 @@ pub trait ExprVisitorMut<'a>: Sized { .chain(&mut window.spec.partition_by) .chain(window.spec.order_by.iter_mut().map(|field| &mut field.expr)) { - self.visit(expr)?; + self.visit(expr, arena)?; } Ok(()) } @@ -139,12 +188,13 @@ pub trait ExprVisitorMut<'a>: Sized { fn visit_in( &mut self, _negated: bool, - expr: &'a mut ScalarExpression, - args: &'a mut [ScalarExpression], + expr: &mut ExprRef, + args: &mut [ExprRef], + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; + self.visit(expr, arena)?; for arg in args { - self.visit(arg)?; + self.visit(arg, arena)?; } Ok(()) } @@ -152,49 +202,53 @@ pub trait ExprVisitorMut<'a>: Sized { fn visit_between( &mut self, _negated: bool, - expr: &'a mut ScalarExpression, - left_expr: &'a mut ScalarExpression, - right_expr: &'a mut ScalarExpression, + expr: &mut ExprRef, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(expr, arena)?; + self.visit(left_expr, arena)?; + self.visit(right_expr, arena) } fn visit_substring( &mut self, - expr: &'a mut ScalarExpression, - for_expr: &'a mut Option>, - from_expr: &'a mut Option>, + expr: &mut ExprRef, + for_expr: &mut Option, + from_expr: &mut Option, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; + self.visit(expr, arena)?; if let Some(for_expr) = for_expr { - self.visit(for_expr)?; + self.visit(for_expr, arena)?; } if let Some(from_expr) = from_expr { - self.visit(from_expr)?; + self.visit(from_expr, arena)?; } Ok(()) } fn visit_position( &mut self, - expr: &'a mut ScalarExpression, - in_expr: &'a mut ScalarExpression, + expr: &mut ExprRef, + in_expr: &mut ExprRef, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; - self.visit(in_expr) + self.visit(expr, arena)?; + self.visit(in_expr, arena) } fn visit_trim( &mut self, - expr: &'a mut ScalarExpression, - trim_what_expr: &'a mut Option>, - _trim_where: &'a mut Option, + expr: &mut ExprRef, + trim_what_expr: &mut Option, + _trim_where: &mut Option, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr)?; + self.visit(expr, arena)?; if let Some(trim_what_expr) = trim_what_expr { - self.visit(trim_what_expr)?; + self.visit(trim_what_expr, arena)?; } Ok(()) } @@ -205,191 +259,205 @@ pub trait ExprVisitorMut<'a>: Sized { fn visit_reference( &mut self, - expr: &'a mut ScalarExpression, + expr: &mut ExprRef, _pos: usize, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } - fn visit_tuple(&mut self, exprs: &'a mut [ScalarExpression]) -> Result<(), DatabaseError> { + fn visit_tuple( + &mut self, + exprs: &mut [ExprRef], + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { for expr in exprs { - self.visit(expr)?; + self.visit(expr, arena)?; } Ok(()) } fn visit_scala_function( &mut self, - scalar_function: &'a mut ScalarFunction, + function: &mut ScalarFunction, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - for arg in &mut scalar_function.args { - self.visit(arg)?; + for arg in &mut function.args { + self.visit(arg, arena)?; } Ok(()) } fn visit_table_function( &mut self, - table_function: &'a mut TableFunction, + function: &mut TableFunction, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - for arg in &mut table_function.args { - self.visit(arg)?; + for arg in &mut function.args { + self.visit(arg, arena)?; } Ok(()) } fn visit_if( &mut self, - condition: &'a mut ScalarExpression, - left_expr: &'a mut ScalarExpression, - right_expr: &'a mut ScalarExpression, - _ty: &'a mut LogicalType, + condition: &mut ExprRef, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(condition)?; - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(condition, arena)?; + self.visit(left_expr, arena)?; + self.visit(right_expr, arena) } fn visit_if_null( &mut self, - left_expr: &'a mut ScalarExpression, - right_expr: &'a mut ScalarExpression, - _ty: &'a mut LogicalType, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left_expr, arena)?; + self.visit(right_expr, arena) } fn visit_null_if( &mut self, - left_expr: &'a mut ScalarExpression, - right_expr: &'a mut ScalarExpression, - _ty: &'a mut LogicalType, + left_expr: &mut ExprRef, + right_expr: &mut ExprRef, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(left_expr)?; - self.visit(right_expr) + self.visit(left_expr, arena)?; + self.visit(right_expr, arena) } fn visit_coalesce( &mut self, - exprs: &'a mut [ScalarExpression], - _ty: &'a mut LogicalType, + exprs: &mut [ExprRef], + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { for expr in exprs { - self.visit(expr)?; + self.visit(expr, arena)?; } Ok(()) } fn visit_case_when( &mut self, - operand_expr: &'a mut Option>, - expr_pairs: &'a mut [(ScalarExpression, ScalarExpression)], - else_expr: &'a mut Option>, - _ty: &'a mut LogicalType, + operand_expr: &mut Option, + expr_pairs: &mut [(ExprRef, ExprRef)], + else_expr: &mut Option, + _ty: &mut LogicalType, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - if let Some(operand_expr) = operand_expr { - self.visit(operand_expr)?; + if let Some(expr) = operand_expr { + self.visit(expr, arena)?; } - for (left_expr, right_expr) in expr_pairs { - self.visit(left_expr)?; - self.visit(right_expr)?; + for (when_expr, then_expr) in expr_pairs { + self.visit(when_expr, arena)?; + self.visit(then_expr, arena)?; } - if let Some(else_expr) = else_expr { - self.visit(else_expr)?; + if let Some(expr) = else_expr { + self.visit(expr, arena)?; } Ok(()) } } -pub fn walk_mut_expr<'a, V: ExprVisitorMut<'a>>( +pub fn walk_mut_expr( visitor: &mut V, - expr: &'a mut ScalarExpression, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - match expr { - ScalarExpression::Constant(value) => visitor.visit_constant(value), + let mut expression = std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + let result = match &mut expression { + ScalarExpression::Constant(value) => visitor.visit_constant(value, arena), ScalarExpression::ColumnRef { column, position } => { - visitor.visit_column_ref(column, position) + visitor.visit_column_ref(column, position, arena) } - ScalarExpression::Alias { expr, alias } => visitor.visit_alias(expr, alias), + ScalarExpression::Alias { expr, alias } => visitor.visit_alias(expr, alias, arena), ScalarExpression::TypeCast { expr, ty, evaluator, - } => visitor.visit_type_cast(expr, ty, evaluator), - ScalarExpression::IsNull { negated, expr } => visitor.visit_is_null(*negated, expr), + } => visitor.visit_type_cast(expr, ty, evaluator, arena), + ScalarExpression::IsNull { negated, expr } => visitor.visit_is_null(*negated, expr, arena), ScalarExpression::Unary { op, expr, evaluator, ty, - } => visitor.visit_unary(op, expr, evaluator, ty), + } => visitor.visit_unary(op, expr, evaluator, ty, arena), ScalarExpression::Binary { op, left_expr, right_expr, evaluator, ty, - } => visitor.visit_binary(op, left_expr, right_expr, evaluator, ty), + } => visitor.visit_binary(op, left_expr, right_expr, evaluator, ty, arena), ScalarExpression::AggCall { distinct, kind, args, ty, - } => visitor.visit_agg(*distinct, kind, args, ty), + } => visitor.visit_agg(*distinct, kind, args, ty, arena), ScalarExpression::In { negated, expr, args, - } => visitor.visit_in(*negated, expr, args), + } => visitor.visit_in(*negated, expr, args, arena), ScalarExpression::Between { negated, expr, left_expr, right_expr, - } => visitor.visit_between(*negated, expr, left_expr, right_expr), + } => visitor.visit_between(*negated, expr, left_expr, right_expr, arena), ScalarExpression::SubString { expr, for_expr, from_expr, - } => visitor.visit_substring(expr, for_expr, from_expr), - ScalarExpression::Position { expr, in_expr } => visitor.visit_position(expr, in_expr), + } => visitor.visit_substring(expr, for_expr, from_expr, arena), + ScalarExpression::Position { expr, in_expr } => { + visitor.visit_position(expr, in_expr, arena) + } ScalarExpression::Trim { expr, trim_what_expr, trim_where, - } => visitor.visit_trim(expr, trim_what_expr, trim_where), + } => visitor.visit_trim(expr, trim_what_expr, trim_where, arena), ScalarExpression::Empty => visitor.visit_empty(), - ScalarExpression::Tuple(exprs) => visitor.visit_tuple(exprs), - ScalarExpression::ScalaFunction(scalar_function) => { - visitor.visit_scala_function(scalar_function) - } - ScalarExpression::TableFunction(table_function) => { - visitor.visit_table_function(table_function) - } + ScalarExpression::Tuple(exprs) => visitor.visit_tuple(exprs, arena), + ScalarExpression::ScalaFunction(function) => visitor.visit_scala_function(function, arena), + ScalarExpression::TableFunction(function) => visitor.visit_table_function(function, arena), ScalarExpression::If { condition, left_expr, right_expr, ty, - } => visitor.visit_if(condition, left_expr, right_expr, ty), + } => visitor.visit_if(condition, left_expr, right_expr, ty, arena), ScalarExpression::IfNull { left_expr, right_expr, ty, - } => visitor.visit_if_null(left_expr, right_expr, ty), + } => visitor.visit_if_null(left_expr, right_expr, ty, arena), ScalarExpression::NullIf { left_expr, right_expr, ty, - } => visitor.visit_null_if(left_expr, right_expr, ty), - ScalarExpression::Coalesce { exprs, ty } => visitor.visit_coalesce(exprs, ty), + } => visitor.visit_null_if(left_expr, right_expr, ty, arena), + ScalarExpression::Coalesce { exprs, ty } => visitor.visit_coalesce(exprs, ty, arena), ScalarExpression::CaseWhen { operand_expr, expr_pairs, else_expr, ty, - } => visitor.visit_case_when(operand_expr, expr_pairs, else_expr, ty), - ScalarExpression::WindowCall(window) => visitor.visit_window(window), - } + } => visitor.visit_case_when(operand_expr, expr_pairs, else_expr, ty, arena), + ScalarExpression::WindowCall(window) => visitor.visit_window(window, arena), + }; + *arena.expression_mut(*expr) = expression; + result } diff --git a/src/expression/window.rs b/src/expression/window.rs index 8c8eed5e..02ba6000 100644 --- a/src/expression/window.rs +++ b/src/expression/window.rs @@ -13,8 +13,8 @@ // limitations under the License. use crate::expression::agg::AggKind; -use crate::expression::ScalarExpression; use crate::planner::operator::sort::SortField; +use crate::planner::ExprRef; use crate::types::LogicalType; use kite_sql_serde_macros::ReferenceSerialization; @@ -49,13 +49,13 @@ impl WindowFunctionKind { #[derive(Debug, Clone, PartialEq, Eq, Hash, ReferenceSerialization)] pub struct WindowFunction { pub kind: WindowFunctionKind, - pub args: Vec, + pub args: Vec, pub ty: LogicalType, } #[derive(Debug, Clone, PartialEq, Eq, Hash, ReferenceSerialization)] pub struct WindowSpec { - pub partition_by: Vec, + pub partition_by: Vec, pub order_by: Vec, } diff --git a/src/function/char_length.rs b/src/function/char_length.rs index 4a5190da..73352dba 100644 --- a/src/function/char_length.rs +++ b/src/function/char_length.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -43,10 +43,11 @@ impl ScalarFunctionImpl for CharLength { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - exprs: &[ScalarExpression], + exprs: &[ExprRef], + arena: &crate::planner::PlanArena<'_>, tuples: Option<&dyn TupleLike>, ) -> Result { - let mut value = exprs[0].eval(tuples)?; + let mut value = arena.expression(exprs[0]).eval(arena, tuples)?; if !matches!(value.logical_type(), LogicalType::Varchar(_, _)) { value = value.cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?; } diff --git a/src/function/current_date.rs b/src/function/current_date.rs index 3c9cd705..f3a5ea30 100644 --- a/src/function/current_date.rs +++ b/src/function/current_date.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -45,7 +45,8 @@ impl ScalarFunctionImpl for CurrentDate { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - _: &[ScalarExpression], + _: &[ExprRef], + _: &crate::planner::PlanArena<'_>, _: Option<&dyn TupleLike>, ) -> Result { Ok(DataValue::Date32(Local::now().num_days_from_ce())) diff --git a/src/function/current_timestamp.rs b/src/function/current_timestamp.rs index 1d7563c3..7c443e8c 100644 --- a/src/function/current_timestamp.rs +++ b/src/function/current_timestamp.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -45,7 +45,8 @@ impl ScalarFunctionImpl for CurrentTimeStamp { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - _: &[ScalarExpression], + _: &[ExprRef], + _: &crate::planner::PlanArena<'_>, _: Option<&dyn TupleLike>, ) -> Result { Ok(DataValue::Time64(Utc::now().timestamp(), 0, false)) diff --git a/src/function/lower.rs b/src/function/lower.rs index c73a1232..c2396395 100644 --- a/src/function/lower.rs +++ b/src/function/lower.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -45,10 +45,11 @@ impl ScalarFunctionImpl for Lower { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - exprs: &[ScalarExpression], + exprs: &[ExprRef], + arena: &crate::planner::PlanArena<'_>, tuples: Option<&dyn TupleLike>, ) -> Result { - let mut value = exprs[0].eval(tuples)?; + let mut value = arena.expression(exprs[0]).eval(arena, tuples)?; if !matches!(value.logical_type(), LogicalType::Varchar(_, _)) { value = value.cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?; } diff --git a/src/function/numbers.rs b/src/function/numbers.rs index 25d48491..08716390 100644 --- a/src/function/numbers.rs +++ b/src/function/numbers.rs @@ -17,8 +17,7 @@ use crate::catalog::ColumnDesc; use crate::errors::DatabaseError; use crate::expression::function::table::TableFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; -use crate::planner::TableArena; +use crate::planner::{ExprRef, TableArena}; use crate::types::tuple::Schema; use crate::types::tuple::Tuple; use crate::types::value::DataValue; @@ -47,9 +46,10 @@ impl TableFunctionImpl for Numbers { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - args: &[ScalarExpression], + args: &[ExprRef], + arena: &crate::planner::PlanArena<'_>, ) -> Result>>, DatabaseError> { - let mut value = args[0].eval::<&Tuple>(None)?; + let mut value = arena.expression(args[0]).eval::<&Tuple>(arena, None)?; value = value.cast(&LogicalType::Integer)?; let num = value diff --git a/src/function/octet_length.rs b/src/function/octet_length.rs index a41b779f..4c31d80d 100644 --- a/src/function/octet_length.rs +++ b/src/function/octet_length.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -44,10 +44,11 @@ impl ScalarFunctionImpl for OctetLength { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - exprs: &[ScalarExpression], + exprs: &[ExprRef], + arena: &crate::planner::PlanArena<'_>, tuples: Option<&dyn TupleLike>, ) -> Result { - let mut value = exprs[0].eval(tuples)?; + let mut value = arena.expression(exprs[0]).eval(arena, tuples)?; if !matches!(value.logical_type(), LogicalType::Varchar(_, _)) { value = value.cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?; } diff --git a/src/function/upper.rs b/src/function/upper.rs index 991896c1..f2a8ac1c 100644 --- a/src/function/upper.rs +++ b/src/function/upper.rs @@ -16,7 +16,7 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; -use crate::expression::ScalarExpression; +use crate::planner::ExprRef; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -45,10 +45,11 @@ impl ScalarFunctionImpl for Upper { #[allow(unused_variables, clippy::redundant_closure_call)] fn eval( &self, - exprs: &[ScalarExpression], + exprs: &[ExprRef], + arena: &crate::planner::PlanArena<'_>, tuples: Option<&dyn TupleLike>, ) -> Result { - let mut value = exprs[0].eval(tuples)?; + let mut value = arena.expression(exprs[0]).eval(arena, tuples)?; if !matches!(value.logical_type(), LogicalType::Varchar(_, _)) { value = value.cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?; } diff --git a/src/macros/mod.rs b/src/macros/mod.rs index dd177b58..e3af1a72 100644 --- a/src/macros/mod.rs +++ b/src/macros/mod.rs @@ -102,11 +102,11 @@ macro_rules! scala_function { } } impl ::kite_sql::expression::function::scala::ScalarFunctionImpl for $struct_name { #[allow(unused_variables, clippy::redundant_closure_call)] - fn eval(&self, args: &[::kite_sql::expression::ScalarExpression], tuple: Option<&dyn ::kite_sql::types::tuple::TupleLike>) -> Result<::kite_sql::types::value::DataValue, ::kite_sql::errors::DatabaseError> { + fn eval(&self, args: &[::kite_sql::planner::ExprRef], arena: &::kite_sql::planner::PlanArena<'_>, tuple: Option<&dyn ::kite_sql::types::tuple::TupleLike>) -> Result<::kite_sql::types::value::DataValue, ::kite_sql::errors::DatabaseError> { let mut _index = 0; $closure($({ - let mut value = args[_index].eval(tuple)?; + let mut value = arena.expression(args[_index]).eval(arena, tuple)?; _index += 1; value = value.cast(&$arg_ty)?; @@ -175,11 +175,11 @@ macro_rules! table_function { impl ::kite_sql::expression::function::table::TableFunctionImpl for $struct_name { #[allow(unused_variables, clippy::redundant_closure_call)] - fn eval(&self, args: &[::kite_sql::expression::ScalarExpression]) -> Result>>, ::kite_sql::errors::DatabaseError> { + fn eval(&self, args: &[::kite_sql::planner::ExprRef], arena: &::kite_sql::planner::PlanArena<'_>) -> Result>>, ::kite_sql::errors::DatabaseError> { let mut _index = 0; $closure($({ - let mut value = args[_index].eval::<&::kite_sql::types::tuple::Tuple>(None)?; + let mut value = arena.expression(args[_index]).eval::<&::kite_sql::types::tuple::Tuple>(arena, None)?; _index += 1; value = value.cast(&$arg_ty)?; diff --git a/src/optimizer/heuristic/optimizer.rs b/src/optimizer/heuristic/optimizer.rs index 06a2cb1a..37189bd7 100644 --- a/src/optimizer/heuristic/optimizer.rs +++ b/src/optimizer/heuristic/optimizer.rs @@ -654,7 +654,7 @@ mod tests { .unwrap(); let mut plan_arena = crate::planner::PlanArena::new(database.state.table_arena()); let sort_fields = vec![SortField::new( - ScalarExpression::column_expr(c1_column, 0), + plan_arena.alloc_expression(ScalarExpression::column_expr(c1_column, 0)), true, false, )]; @@ -724,13 +724,23 @@ mod tests { } ]))) ); - assert_eq!( - physical_option.sort_option(), - &SortOption::OrderBy { - fields: sort_fields, - ignore_prefix_len: 0, - } - ); + let SortOption::OrderBy { + fields, + ignore_prefix_len, + } = physical_option.sort_option() + else { + panic!("index scan should preserve the requested order") + }; + assert_eq!(*ignore_prefix_len, 0); + assert_eq!(fields.len(), sort_fields.len()); + for (actual, expected) in fields.iter().zip(&sort_fields) { + assert_eq!(actual.asc, expected.asc); + assert_eq!(actual.nulls_first, expected.nulls_first); + assert_eq!( + plan_arena.expression(actual.expr), + plan_arena.expression(expected.expr) + ); + } Ok(()) } diff --git a/src/optimizer/rule/implementation/mod.rs b/src/optimizer/rule/implementation/mod.rs index 97855b82..6f9fc067 100644 --- a/src/optimizer/rule/implementation/mod.rs +++ b/src/optimizer/rule/implementation/mod.rs @@ -314,6 +314,7 @@ mod tests { use crate::expression::function::table::{ ArcTableFunctionImpl, TableFunction, TableFunctionCatalog, TableFunctionImpl, }; + use crate::expression::ScalarExpression; use crate::function::numbers::Numbers; use crate::optimizer::core::rule::{ImplementationRule, MatchPattern}; use crate::optimizer::core::statistics_meta::StatisticMetaLoader; @@ -521,35 +522,31 @@ mod tests { } if fields.len() == 1 )); - let function_operator = function_scan_operator(); - let function_option = best_option( - ImplementationRuleImpl::FunctionScan, - &function_operator, - arena, - )?; - assert_eq!(function_option.plan, PlanImpl::FunctionScan); - Ok(()) }, )?; - - Ok(()) - } - - fn function_scan_operator() -> Operator { let table_arena = TableArenaCell::default(); let numbers = Numbers::new(); let mut schema = Vec::new(); numbers.output_schema_into(table_arena.borrow_mut(), &mut schema); - let table_function = TableFunction { - args: vec![DataValue::Int32(3).into()], - catalog: TableFunctionCatalog { - schema, - inner: ArcTableFunctionImpl(numbers), + let mut arena = PlanArena::new(&table_arena); + let function_operator = Operator::FunctionScan(FunctionScanOperator { + table_function: TableFunction { + args: vec![arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(3)))], + catalog: TableFunctionCatalog { + schema, + inner: ArcTableFunctionImpl(numbers), + }, }, - }; + }); + let function_option = best_option( + ImplementationRuleImpl::FunctionScan, + &function_operator, + &arena, + )?; + assert_eq!(function_option.plan, PlanImpl::FunctionScan); - Operator::FunctionScan(FunctionScanOperator { table_function }) + Ok(()) } #[test] diff --git a/src/optimizer/rule/normalization/column_pruning.rs b/src/optimizer/rule/normalization/column_pruning.rs index 3c6a8e24..91123f0c 100644 --- a/src/optimizer/rule/normalization/column_pruning.rs +++ b/src/optimizer/rule/normalization/column_pruning.rs @@ -25,10 +25,11 @@ use crate::planner::operator::join::JoinCondition; use crate::planner::operator::visitor::{OperatorExprVisitor, OperatorVisitor}; use crate::planner::operator::visitor_mut::{OperatorExprVisitorMut, OperatorVisitorMut}; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; use crate::types::LogicalType; +use std::collections::HashSet; #[derive(Clone)] pub struct ColumnPruning; @@ -36,6 +37,7 @@ pub struct ColumnPruning; struct ApplyOutcome { changed: bool, removed_positions: Vec, + remapped_exprs: HashSet, } #[derive(Clone, Default)] @@ -88,7 +90,7 @@ struct ReferencedColumnCollector<'a, 'p> { arena: &'a crate::planner::PlanArena<'p>, } -impl ExprVisitor<'_> for ReferencedColumnCollector<'_, '_> { +impl ExprVisitor> for ReferencedColumnCollector<'_, '_> { fn visit_column_ref( &mut self, column: &crate::catalog::ColumnRef, @@ -99,10 +101,11 @@ impl ExprVisitor<'_> for ReferencedColumnCollector<'_, '_> { fn visit_alias( &mut self, - expr: &ScalarExpression, + expr: ExprRef, _ty: &AliasType, + arena: &PlanArena<'_>, ) -> Result<(), DatabaseError> { - self.visit(expr) + self.visit(expr, arena) } } @@ -111,6 +114,7 @@ impl ApplyOutcome { Self { changed: false, removed_positions: Vec::with_capacity(arena.allocated_columns_len()), + remapped_exprs: HashSet::new(), } } } @@ -125,7 +129,7 @@ impl ColumnPruning { referenced_columns, arena, }; - OperatorExprVisitor::new(&mut collector).visit_operator(operator)?; + OperatorExprVisitor::new(&mut collector, arena).visit_operator(operator)?; struct ReferencedOperatorColumnCollector<'a, 'p> { referenced_columns: &'a mut ReferencedColumns, @@ -206,7 +210,7 @@ impl ColumnPruning { } fn extend_expr_referenced_columns<'a>( - exprs: impl IntoIterator, + exprs: impl IntoIterator, referenced_columns: &mut ReferencedColumns, arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { @@ -215,13 +219,13 @@ impl ColumnPruning { arena, }; for expr in exprs { - collector.visit(expr)?; + collector.visit(*expr, collector.arena)?; } Ok(()) } fn output_column_is_required( - expr: &ScalarExpression, + expr: ExprRef, column_references: &ReferencedColumns, arena: &mut crate::planner::PlanArena, ) -> bool { @@ -231,7 +235,7 @@ impl ColumnPruning { fn clear_exprs( column_references: &ReferencedColumns, - exprs: &mut Vec, + exprs: &mut Vec, removed_positions: &mut Vec, output_start: usize, arena: &mut crate::planner::PlanArena, @@ -240,7 +244,7 @@ impl ColumnPruning { removed_positions.reserve(exprs.len()); let mut position = 0; exprs.retain(|expr| { - let keep = Self::output_column_is_required(expr, column_references, arena); + let keep = Self::output_column_is_required(*expr, column_references, arena); if !keep { removed_positions.push(position); } @@ -252,19 +256,26 @@ impl ColumnPruning { fn remap_operator_after_child_change( operator: &mut Operator, removed_positions: &[usize], + remapped_exprs: &mut HashSet, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { - OperatorExprVisitorMut::new(&mut PositionRemapper { removed_positions }) - .visit_operator(operator) + OperatorExprVisitorMut::new( + &mut PositionRemapper::new(removed_positions, remapped_exprs), + arena, + ) + .visit_operator(operator) } fn remap_exprs_after_child_change<'a>( - exprs: impl IntoIterator, + exprs: impl IntoIterator, removed_positions: &[usize], + remapped_exprs: &mut HashSet, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { if removed_positions.is_empty() { return Ok(()); } - remap_exprs_positions(exprs, removed_positions) + remap_exprs_positions(exprs, removed_positions, remapped_exprs, arena) } fn apply_only_child( @@ -379,12 +390,14 @@ impl ColumnPruning { }; // only single COUNT(*) is not depend on any column // removed all expressions from the aggregate: push a COUNT(*) - op.agg_calls.push(ScalarExpression::AggCall { - distinct: false, - kind: AggKind::Count, - args: vec![ScalarExpression::Constant(value)], - ty: LogicalType::Integer, - }); + let arg = arena.alloc_expression(ScalarExpression::Constant(value)); + op.agg_calls + .push(arena.alloc_expression(ScalarExpression::AggCall { + distinct: false, + kind: AggKind::Count, + args: vec![arg], + ty: LogicalType::Integer, + })); changed = true; } } else { @@ -415,6 +428,8 @@ impl ColumnPruning { Self::remap_operator_after_child_change( operator, &outcome.removed_positions[child_start..], + &mut outcome.remapped_exprs, + arena, )?; changed = true; } @@ -423,7 +438,7 @@ impl ColumnPruning { Operator::Project(op) => { let mut has_count_star = HasCountStar::default(); for expr in &op.exprs { - has_count_star.visit(expr)?; + has_count_star.visit(*expr, arena)?; } if !has_count_star.value { if !all_referenced { @@ -463,6 +478,8 @@ impl ColumnPruning { Self::remap_operator_after_child_change( operator, &outcome.removed_positions[child_start..], + &mut outcome.remapped_exprs, + arena, )?; changed = true; } @@ -568,16 +585,18 @@ impl ColumnPruning { outcome.removed_positions [left_removed_start..right_removed_end] .split_at(left_removed_len); - for (left_expr, right_expr) in on { - remap_expr_positions( - left_expr, - left_removed_positions, - )?; - remap_expr_positions( - right_expr, - right_removed_positions, - )?; - } + remap_exprs_positions( + on.iter_mut().map(|(left_expr, _)| left_expr), + left_removed_positions, + &mut outcome.remapped_exprs, + arena, + )?; + remap_exprs_positions( + on.iter_mut().map(|(_, right_expr)| right_expr), + right_removed_positions, + &mut outcome.remapped_exprs, + arena, + )?; } Self::offset_removed_positions( &mut outcome.removed_positions @@ -588,7 +607,12 @@ impl ColumnPruning { let removed_positions = &outcome.removed_positions [left_removed_start..right_removed_end]; if !removed_positions.is_empty() { - remap_expr_positions(filter, removed_positions)?; + remap_expr_positions( + *filter, + removed_positions, + &mut outcome.remapped_exprs, + arena, + )?; } } } @@ -611,6 +635,8 @@ impl ColumnPruning { Self::remap_exprs_after_child_change( op.predicates_mut().iter_mut(), removed_positions, + &mut outcome.remapped_exprs, + arena, )?; outcome.removed_positions.truncate(right_removed_start); } else { @@ -657,6 +683,8 @@ impl ColumnPruning { Self::remap_operator_after_child_change( operator, &outcome.removed_positions[child_start..], + &mut outcome.remapped_exprs, + arena, )?; changed = true; } @@ -694,6 +722,8 @@ impl ColumnPruning { Self::remap_operator_after_child_change( operator, &outcome.removed_positions[child_start..], + &mut outcome.remapped_exprs, + arena, )?; changed = true; } @@ -726,6 +756,8 @@ impl ColumnPruning { Self::remap_operator_after_child_change( operator, &outcome.removed_positions[child_start..], + &mut outcome.remapped_exprs, + arena, )?; changed = true; } diff --git a/src/optimizer/rule/normalization/combine_operators.rs b/src/optimizer/rule/normalization/combine_operators.rs index 7c7e2712..dd28d35d 100644 --- a/src/optimizer/rule/normalization/combine_operators.rs +++ b/src/optimizer/rule/normalization/combine_operators.rs @@ -13,89 +13,101 @@ // limitations under the License. use crate::errors::DatabaseError; -use crate::expression::{AliasType, BinaryOperator, ScalarExpression}; +use crate::expression::visitor_mut::ExprVisitorMut; +use crate::expression::{BinaryOperator, ScalarExpression}; use crate::optimizer::core::rule::NormalizationRule; use crate::optimizer::plan_utils::{only_child_mut, replace_with_only_child}; use crate::optimizer::rule::normalization::strip_alias; use crate::planner::operator::filter::FilterOperator; use crate::planner::operator::project::ProjectOperator; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; -use crate::types::value::DataValue; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::LogicalType; use std::mem; -fn is_passthrough_project(op: &ProjectOperator) -> bool { - op.exprs - .iter() - .all(|expr| matches!(strip_alias(expr), ScalarExpression::ColumnRef { .. })) +fn is_passthrough_project(op: &ProjectOperator, arena: &PlanArena<'_>) -> bool { + op.exprs.iter().all(|expr| { + matches!( + arena.expression(strip_alias(*expr, arena)), + ScalarExpression::ColumnRef { .. } + ) + }) } -fn passthrough_source_position(expr: &ScalarExpression) -> Option { - match strip_alias(expr) { +fn passthrough_source_position(expr: ExprRef, arena: &PlanArena<'_>) -> Option { + match arena.expression(strip_alias(expr, arena)) { ScalarExpression::ColumnRef { position, .. } => Some(*position), _ => None, } } -fn collapse_match_expr(expr: &ScalarExpression) -> &ScalarExpression { - match expr { - ScalarExpression::Alias { expr, .. } => expr, +fn collapse_match_expr(expr: ExprRef, arena: &PlanArena<'_>) -> ExprRef { + match arena.expression(expr) { + ScalarExpression::Alias { expr, .. } => *expr, _ => expr, } } -fn rewrite_column_position(expr: &mut ScalarExpression, new_position: usize) { - match expr { - ScalarExpression::ColumnRef { position, .. } => { - *position = new_position; +fn rewrite_column_position( + mut expr: ExprRef, + new_position: usize, + arena: &mut PlanArena<'_>, +) -> Result<(), DatabaseError> { + struct Rewriter(usize); + impl ExprVisitorMut for Rewriter { + fn visit_column_ref( + &mut self, + _column: &mut crate::catalog::ColumnRef, + position: &mut usize, + _arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + *position = self.0; + Ok(()) } - ScalarExpression::Alias { expr, alias } => { - rewrite_column_position(expr, new_position); - if let AliasType::Expr(alias_expr) = alias { - rewrite_column_position(alias_expr, new_position); - } - } - _ => {} } + Rewriter(new_position).visit(&mut expr, arena) } fn remap_passthrough_project_exprs( - parent_exprs: &mut [ScalarExpression], - child_exprs: &[ScalarExpression], - arena: &crate::planner::PlanArena, -) -> bool { - let mut remapped_positions = Vec::with_capacity(parent_exprs.len()); + parent_exprs: &mut [ExprRef], + child_exprs: &[ExprRef], + remapped_positions: &mut Vec, + arena: &mut PlanArena<'_>, +) -> Result { + remapped_positions.clear(); for parent_expr in parent_exprs.iter() { - let parent_match_expr = collapse_match_expr(parent_expr); + let parent_match_expr = collapse_match_expr(*parent_expr, arena); let Some(position) = child_exprs .iter() - .find(|child_expr| parent_match_expr.eq_ignore_colref_pos(child_expr, arena)) - .and_then(passthrough_source_position) + .find(|child_expr| parent_match_expr.eq_ignore_colref_pos(**child_expr, arena)) + .and_then(|expr| passthrough_source_position(*expr, arena)) else { - return false; + return Ok(false); }; remapped_positions.push(position); } - for (parent_expr, position) in parent_exprs.iter_mut().zip(remapped_positions) { - rewrite_column_position(parent_expr, position); + for (parent_expr, position) in parent_exprs + .iter_mut() + .zip(remapped_positions.iter().copied()) + { + rewrite_column_position(*parent_expr, position, arena)?; } - true + Ok(true) } fn groupby_exprs_match( - parent_exprs: &[ScalarExpression], - child_exprs: &[ScalarExpression], + parent_exprs: &[ExprRef], + child_exprs: &[ExprRef], arena: &crate::planner::PlanArena, ) -> bool { parent_exprs.len() == child_exprs.len() && parent_exprs .iter() .zip(child_exprs.iter()) - .all(|(parent_expr, child_expr)| parent_expr.eq_ignore_colref_pos(child_expr, arena)) + .all(|(parent_expr, child_expr)| parent_expr.eq_ignore_colref_pos(*child_expr, arena)) } /// Combine two adjacent project operators into one. @@ -112,18 +124,20 @@ impl NormalizationRule for CollapseProject { }; let mut removed = false; + let mut remapped_positions = Vec::with_capacity(parent_op.exprs.len()); loop { let Childrens::Only(child) = plan.childrens.as_mut() else { break; }; match &child.operator { Operator::Project(child_op) - if is_passthrough_project(child_op) + if is_passthrough_project(child_op, arena) && remap_passthrough_project_exprs( &mut parent_op.exprs, &child_op.exprs, + &mut remapped_positions, arena, - ) => + )? => { removed |= replace_with_only_child(child.as_mut()); } @@ -142,7 +156,7 @@ impl NormalizationRule for CombineFilter { fn apply( &self, plan: &mut LogicalPlan, - _: &mut crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result { let parent_filter = match mem::replace(&mut plan.operator, Operator::Dummy) { Operator::Filter(op) => op, @@ -169,22 +183,19 @@ impl NormalizationRule for CombineFilter { having, is_optimized: _, } = parent_filter; - let child_predicate = mem::replace( - &mut child_op.predicate, - ScalarExpression::Constant(DataValue::Boolean(true)), - ); - child_op.predicate = ScalarExpression::Binary { + let child_predicate = child_op.predicate; + child_op.predicate = arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(predicate), - right_expr: Box::new(child_predicate), + left_expr: predicate, + right_expr: child_predicate, evaluator: None, ty: LogicalType::Boolean, - }; + }); child_op.having = having || child_op.having; return Ok(replace_with_only_child(plan)); } - Operator::Project(project_op) if is_passthrough_project(project_op) => { + Operator::Project(project_op) if is_passthrough_project(project_op, arena) => { if replace_with_only_child(cursor) { continue; } @@ -254,11 +265,11 @@ mod tests { use crate::planner::operator::aggregate::AggregateOperator; use crate::planner::operator::project::ProjectOperator; use crate::planner::operator::Operator; - use crate::planner::{Childrens, LogicalPlan, PlanArena}; + use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; - fn column_expr(arena: &mut PlanArena, name: &str, position: usize) -> ScalarExpression { + fn column_expr(arena: &mut PlanArena, name: &str, position: usize) -> ExprRef { let column = arena.alloc_column(ColumnCatalog::new_dummy(name.to_string())); - ScalarExpression::column_expr(column, position) + arena.alloc_expression(ScalarExpression::column_expr(column, position)) } #[test] @@ -357,7 +368,7 @@ mod tests { let Operator::Project(op) = &plan.operator else { unreachable!("expected project"); }; - let ScalarExpression::ColumnRef { position, .. } = &op.exprs[0] else { + let ScalarExpression::ColumnRef { position, .. } = arena.expression(op.exprs[0]) else { unreachable!("expected column ref"); }; assert_eq!(*position, 1); @@ -388,7 +399,7 @@ mod tests { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(op) = &filter_op.operator { - if let ScalarExpression::Binary { op, .. } = &op.predicate { + if let ScalarExpression::Binary { op, .. } = arena.expression(op.predicate) { assert_eq!(op, &BinaryOperator::And); } else { unreachable!("Should be a and operator") @@ -447,7 +458,8 @@ mod tests { let Operator::Aggregate(op) = &plan.operator else { unreachable!("expected aggregate"); }; - let ScalarExpression::ColumnRef { position, .. } = &op.groupby_exprs[0] else { + let ScalarExpression::ColumnRef { position, .. } = arena.expression(op.groupby_exprs[0]) + else { unreachable!("expected column ref"); }; assert_eq!(*position, 1); diff --git a/src/optimizer/rule/normalization/compilation_in_advance.rs b/src/optimizer/rule/normalization/compilation_in_advance.rs index d0aa1bd0..08a909cc 100644 --- a/src/optimizer/rule/normalization/compilation_in_advance.rs +++ b/src/optimizer/rule/normalization/compilation_in_advance.rs @@ -24,14 +24,14 @@ pub struct EvaluatorBind; pub(crate) fn evaluator_bind_current( plan: &mut LogicalPlan, - arena: &PlanArena, + arena: &mut PlanArena, ) -> Result<(), DatabaseError> { - let mut evaluator = BindEvaluator { arena }; - OperatorExprVisitorMut::new(&mut evaluator).visit_operator(&mut plan.operator) + let mut evaluator = BindEvaluator; + OperatorExprVisitorMut::new(&mut evaluator, arena).visit_operator(&mut plan.operator) } impl EvaluatorBind { - fn _apply(plan: &mut LogicalPlan, arena: &PlanArena) -> Result<(), DatabaseError> { + fn _apply(plan: &mut LogicalPlan, arena: &mut PlanArena) -> Result<(), DatabaseError> { match plan.childrens.as_mut() { Childrens::Only(child) => Self::_apply(child, arena)?, Childrens::Twins { left, right } => { @@ -86,23 +86,25 @@ mod tests { use crate::planner::operator::union::UnionOperator; use crate::planner::operator::update::UpdateOperator; use crate::planner::operator::Operator; - use crate::planner::TableArenaCell; + use crate::planner::{ExprRef, TableArenaCell}; use crate::types::value::DataValue; use crate::types::LogicalType; - fn unbound_binary(left: ScalarExpression, right: ScalarExpression) -> ScalarExpression { - ScalarExpression::Binary { + fn unbound_binary(arena: &mut PlanArena, column: crate::catalog::ColumnRef) -> ExprRef { + let left_expr = arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let right_expr = arena.alloc_expression(DataValue::Int32(1).into()); + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Plus, - left_expr: Box::new(left), - right_expr: Box::new(right), + left_expr, + right_expr, evaluator: None, ty: LogicalType::Integer, - } + }) } - fn is_bound(expr: &ScalarExpression) -> bool { + fn is_bound(expr: ExprRef, arena: &PlanArena) -> bool { matches!( - expr, + arena.expression(expr), ScalarExpression::Binary { evaluator: Some(_), .. @@ -122,37 +124,37 @@ mod tests { false, ColumnDesc::new(LogicalType::Integer, None, false, None)?, )); - let expr = || { - unbound_binary( - ScalarExpression::column_expr(column, 0), - DataValue::Int32(1).into(), - ) - }; + let filter_expr = unbound_binary(&mut arena, column); + let sort_expr = unbound_binary(&mut arena, column); + let top_k_expr = unbound_binary(&mut arena, column); + let mark_expr = unbound_binary(&mut arena, column); + let update_expr = unbound_binary(&mut arena, column); + let function_expr = unbound_binary(&mut arena, column); let mut operators = vec![ Operator::Filter(FilterOperator { - predicate: expr(), + predicate: filter_expr, is_optimized: false, having: false, }), Operator::Sort(SortOperator { - sort_fields: vec![SortField::from(expr())], + sort_fields: vec![SortField::from(sort_expr)], }), Operator::TopK(TopKOperator { - sort_fields: vec![SortField::from(expr())], + sort_fields: vec![SortField::from(top_k_expr)], limit: 1, offset: None, }), - Operator::MarkApply(MarkApplyOperator::new_exists(column, vec![expr()])), + Operator::MarkApply(MarkApplyOperator::new_exists(column, vec![mark_expr])), Operator::Update(UpdateOperator { table_name: "t1".into(), - value_exprs: vec![(column, expr())], + value_exprs: vec![(column, update_expr)], }), ]; operators.push(Operator::FunctionScan(FunctionScanOperator { table_function: TableFunction { - args: vec![expr()], + args: vec![function_expr], catalog: TableFunctionCatalog { schema, inner: ArcTableFunctionImpl(numbers), @@ -162,14 +164,14 @@ mod tests { for operator in &mut operators { let mut plan = LogicalPlan::new(operator.clone(), Childrens::None); - evaluator_bind_current(&mut plan, &arena)?; + evaluator_bind_current(&mut plan, &mut arena)?; match &plan.operator { - Operator::Filter(op) => assert!(is_bound(&op.predicate)), - Operator::Sort(op) => assert!(is_bound(&op.sort_fields[0].expr)), - Operator::TopK(op) => assert!(is_bound(&op.sort_fields[0].expr)), - Operator::MarkApply(op) => assert!(is_bound(&op.predicates()[0])), - Operator::Update(op) => assert!(is_bound(&op.value_exprs[0].1)), - Operator::FunctionScan(op) => assert!(is_bound(&op.table_function.args[0])), + Operator::Filter(op) => assert!(is_bound(op.predicate, &arena)), + Operator::Sort(op) => assert!(is_bound(op.sort_fields[0].expr, &arena)), + Operator::TopK(op) => assert!(is_bound(op.sort_fields[0].expr, &arena)), + Operator::MarkApply(op) => assert!(is_bound(op.predicates()[0], &arena)), + Operator::Update(op) => assert!(is_bound(op.value_exprs[0].1, &arena)), + Operator::FunctionScan(op) => assert!(is_bound(op.table_function.args[0], &arena)), _ => unreachable!(), } } @@ -186,35 +188,36 @@ mod tests { false, ColumnDesc::new(LogicalType::Integer, None, false, None)?, )); - let expr = || { - unbound_binary( - ScalarExpression::column_expr(column, 0), - DataValue::Int32(1).into(), - ) - }; - let filter = || { + fn filter(arena: &mut PlanArena, column: crate::catalog::ColumnRef) -> LogicalPlan { + let predicate = unbound_binary(arena, column); LogicalPlan::new( Operator::Filter(FilterOperator { - predicate: expr(), + predicate, is_optimized: false, having: false, }), Childrens::None, ) - }; + } + + let join_left = unbound_binary(&mut arena, column); + let join_right = unbound_binary(&mut arena, column); + let join_filter = unbound_binary(&mut arena, column); + let left_filter = filter(&mut arena, column); + let right_filter = filter(&mut arena, column); let mut join = LogicalPlan::new( Operator::Join(JoinOperator { join_type: JoinType::Inner, force_nested_loop: false, on: JoinCondition::On { - on: vec![(expr(), expr())], - filter: Some(expr()), + on: vec![(join_left, join_right)], + filter: Some(join_filter), }, }), Childrens::Twins { - left: Box::new(filter()), - right: Box::new(filter()), + left: Box::new(left_filter), + right: Box::new(right_filter), }, ); assert!(EvaluatorBind.apply(&mut join, &mut arena)?); @@ -228,17 +231,19 @@ mod tests { else { unreachable!() }; - assert!(is_bound(&on[0].0)); - assert!(is_bound(&on[0].1)); - assert!(is_bound(join_filter.as_ref().unwrap())); + assert!(is_bound(on[0].0, &arena)); + assert!(is_bound(on[0].1, &arena)); + assert!(is_bound(*join_filter.as_ref().unwrap(), &arena)); assert!(join.childrens.iter().all(|child| { - matches!(&child.operator, Operator::Filter(op) if is_bound(&op.predicate)) + matches!(&child.operator, Operator::Filter(op) if is_bound(op.predicate, &arena)) })); - let mut union = UnionOperator::build(vec![column], vec![column], filter(), filter()); + let union_left = filter(&mut arena, column); + let union_right = filter(&mut arena, column); + let mut union = UnionOperator::build(vec![column], vec![column], union_left, union_right); EvaluatorBind.apply(&mut union, &mut arena)?; assert!(union.childrens.iter().all(|child| { - matches!(&child.operator, Operator::Filter(op) if is_bound(&op.predicate)) + matches!(&child.operator, Operator::Filter(op) if is_bound(op.predicate, &arena)) })); Ok(()) diff --git a/src/optimizer/rule/normalization/elimination.rs b/src/optimizer/rule/normalization/elimination.rs index f1bde9a1..0517c588 100644 --- a/src/optimizer/rule/normalization/elimination.rs +++ b/src/optimizer/rule/normalization/elimination.rs @@ -13,14 +13,13 @@ // limitations under the License. use crate::errors::DatabaseError; -use crate::expression::ScalarExpression; use crate::optimizer::core::rule::NormalizationRule; use crate::optimizer::plan_utils::{only_child_mut, replace_with_only_child, wrap_child_with}; use crate::planner::operator::limit::LimitOperator; use crate::planner::operator::sort::{SortField, SortOperator}; use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::{Operator, PhysicalOption, PlanImpl, SortOption}; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::types::index::{IndexLookup, IndexOrderHint}; pub struct EliminateRedundantSort; @@ -90,7 +89,7 @@ impl NormalizationRule for EliminateIndexFilter { if !matches!(index_info.lookup, Some(IndexLookup::Static(_))) { return Ok(false); } - index_info.residual_predicate.clone() + index_info.residual_predicate }; if let Some(residual) = residual { @@ -122,7 +121,7 @@ pub(crate) enum OrderHintKind { #[derive(Copy, Clone)] pub(crate) enum ScanOrderHint<'a> { SortFields(&'a [SortField]), - GroupBy(&'a [ScalarExpression]), + GroupBy(&'a [ExprRef]), } impl<'a> ScanOrderHint<'a> { @@ -130,7 +129,7 @@ impl<'a> ScanOrderHint<'a> { Self::SortFields(fields) } - pub(crate) fn groupby(groupby_exprs: &'a [ScalarExpression]) -> Self { + pub(crate) fn groupby(groupby_exprs: &'a [ExprRef]) -> Self { Self::GroupBy(groupby_exprs) } } @@ -172,8 +171,8 @@ pub(crate) fn apply_scan_order_hint( let mut required_from_table = true; for index in 0..hint_len(required) { let expr = match required { - ScanOrderHint::SortFields(fields) => &fields[index].expr, - ScanOrderHint::GroupBy(groupby_exprs) => &groupby_exprs[index], + ScanOrderHint::SortFields(fields) => fields[index].expr, + ScanOrderHint::GroupBy(groupby_exprs) => groupby_exprs[index], }; if !expr.all_referenced_columns(arena, |arena, column| { scan_op @@ -229,15 +228,15 @@ fn hint_covers( sort_field_matches(required, provided, arena) }), ScanOrderHint::GroupBy(groupby_exprs) => covers(groupby_exprs, provided, |expr, field| { - field.asc && !field.nulls_first && expr.eq_ignore_colref_pos(&field.expr, arena) + field.asc && !field.nulls_first && expr.eq_ignore_colref_pos(field.expr, arena) }), } } -pub(crate) fn groupby_sort_fields(groupby_exprs: &[ScalarExpression]) -> Vec { +pub(crate) fn groupby_sort_fields(groupby_exprs: &[ExprRef]) -> Vec { groupby_exprs .iter() - .cloned() + .copied() .map(|expr| SortField::new(expr, true, false)) .collect() } @@ -401,7 +400,7 @@ fn sort_field_matches( ) -> bool { required.asc == provided.asc && required.nulls_first == provided.nulls_first - && required.expr.eq_ignore_colref_pos(&provided.expr, arena) + && required.expr.eq_ignore_colref_pos(provided.expr, arena) } pub(crate) fn covers( @@ -457,7 +456,7 @@ mod tests { use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::top_k::TopKOperator; use crate::planner::operator::{Operator, PhysicalOption, PlanImpl, SortOption}; - use crate::planner::{Childrens, LogicalPlan}; + use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::types::index::{IndexInfo, IndexLookup, IndexMeta, IndexType}; use crate::types::value::DataValue; use crate::types::ColumnId; @@ -473,7 +472,11 @@ mod tests { position: usize, ) -> SortField { let column = arena.alloc_column(ColumnCatalog::new_dummy(name.to_string())); - SortField::new(ScalarExpression::column_expr(column, position), true, false) + SortField::new( + arena.alloc_expression(ScalarExpression::column_expr(column, position)), + true, + false, + ) } fn build_plan( @@ -491,9 +494,11 @@ mod tests { index_sort_option, )); + let predicate = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); let mut filter = LogicalPlan::new( Operator::Filter(FilterOperator { - predicate: ScalarExpression::Constant(DataValue::Boolean(true)), + predicate, is_optimized: false, having: false, }), @@ -546,11 +551,15 @@ mod tests { fn build_filter_with_selected_index( arena: &mut crate::planner::PlanArena, - predicate: ScalarExpression, - residual: Option, + predicate: ExprRef, + residual: Option, ) -> LogicalPlan { let column = arena.alloc_column(ColumnCatalog::new_dummy("c1".to_string())); - let sort_field = SortField::new(ScalarExpression::column_expr(column, 0), true, false); + let sort_field = SortField::new( + arena.alloc_expression(ScalarExpression::column_expr(column, 0)), + true, + false, + ); let (mut index_info, sort_option) = build_index_info(arena, vec![sort_field], 0); index_info.lookup = Some(IndexLookup::Static(Range::Scope { min: Bound::Unbounded, @@ -593,13 +602,10 @@ mod tests { let c1_id = 1; let columns = vec![c1]; - let sort_fields = vec![SortField::new( - ScalarExpression::column_expr(c1, 0), - true, - false, - )]; + let group_expr = arena.alloc_expression(ScalarExpression::column_expr(c1, 0)); + let sort_fields = vec![SortField::new(group_expr, true, false)]; let sort_option = SortOption::OrderBy { - fields: sort_fields.clone(), + fields: sort_fields, ignore_prefix_len: 0, }; let index_info = IndexInfo { @@ -634,7 +640,7 @@ mod tests { let plan = LogicalPlan::new( Operator::Aggregate(AggregateOperator { - groupby_exprs: vec![ScalarExpression::column_expr(c1, 0)], + groupby_exprs: vec![group_expr], agg_calls: vec![], is_distinct: true, force_spill: false, @@ -649,7 +655,8 @@ mod tests { fn exact_index_filter_is_removed_after_physical_selection() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = crate::planner::PlanArena::new(&table_arena); - let predicate = ScalarExpression::Constant(DataValue::Boolean(true)); + let predicate = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); let mut plan = build_filter_with_selected_index(&mut arena, predicate, None); let rule = EliminateIndexFilter; @@ -662,10 +669,11 @@ mod tests { fn partial_index_filter_keeps_residual_after_physical_selection() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = crate::planner::PlanArena::new(&table_arena); - let predicate = ScalarExpression::Constant(DataValue::Boolean(true)); - let residual = ScalarExpression::Constant(DataValue::Boolean(false)); - let mut plan = - build_filter_with_selected_index(&mut arena, predicate, Some(residual.clone())); + let predicate = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); + let residual = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(false))); + let mut plan = build_filter_with_selected_index(&mut arena, predicate, Some(residual)); let rule = EliminateIndexFilter; assert!(rule.apply(&mut plan, &mut arena)?); @@ -680,7 +688,8 @@ mod tests { fn probe_index_filter_is_not_removed_after_physical_selection() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = crate::planner::PlanArena::new(&table_arena); - let predicate = ScalarExpression::Constant(DataValue::Boolean(true)); + let predicate = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); let mut plan = build_filter_with_selected_index(&mut arena, predicate, None); let Childrens::Only(child) = plan.childrens.as_mut() else { unreachable!("filter should have a scan child"); @@ -792,14 +801,18 @@ mod tests { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = crate::planner::PlanArena::new(&table_arena); let column = arena.alloc_column(ColumnCatalog::new_dummy("c1".to_string())); - let sort_field = SortField::new(ScalarExpression::column_expr(column, 0), true, false); + let sort_field = SortField::new( + arena.alloc_expression(ScalarExpression::column_expr(column, 0)), + true, + false, + ); let (index_info, _) = build_index_info(&mut arena, vec![sort_field.clone()], 0); let columns = vec![column]; let table_name: TableName = ::std::sync::Arc::from("t"); let table_scan = LogicalPlan::new( Operator::TableScan(TableScanOperator { - table_name: table_name.clone(), + table_name, columns, limit: (None, None), index_infos: vec![index_info], @@ -878,7 +891,7 @@ mod tests { let index_info = scan_op.index_infos[0].clone(); child.physical_option = Some(PhysicalOption::new( PlanImpl::IndexScan(Box::new(index_info)), - sort_option.clone(), + sort_option, )); } } @@ -1012,12 +1025,7 @@ mod tests { let mut arena = crate::planner::PlanArena::new(&table_arena); let c1 = make_sort_field(&mut arena, "c1"); let c2 = make_sort_field(&mut arena, "c2"); - let mut plan = build_plan( - &mut arena, - vec![c2.clone()], - vec![c1.clone(), c2.clone()], - 0, - ); + let mut plan = build_plan(&mut arena, vec![c2.clone()], vec![c1, c2.clone()], 0); super::mark_sort_preserving_indexes(&mut plan, &[c2], &arena)?; let rule = EliminateRedundantSort; @@ -1031,7 +1039,11 @@ mod tests { let table_arena = crate::planner::TableArenaCell::default(); let mut arena = crate::planner::PlanArena::new(&table_arena); let column = arena.alloc_column(ColumnCatalog::new_dummy("c_first".to_string())); - let sort_field = SortField::new(ScalarExpression::column_expr(column, 0), true, false); + let sort_field = SortField::new( + arena.alloc_expression(ScalarExpression::column_expr(column, 0)), + true, + false, + ); let (mut index_info, _) = build_index_info(&mut arena, vec![sort_field.clone()], 0); index_info.lookup = Some(IndexLookup::Static(Range::Scope { min: Bound::Unbounded, @@ -1054,13 +1066,14 @@ mod tests { let index_info = scan_op.index_infos[0].clone(); scan_plan.physical_option = Some(PhysicalOption::new( PlanImpl::IndexScan(Box::new(index_info.clone())), - index_info.sort_option.clone(), + index_info.sort_option, )); } let mut filter = LogicalPlan::new( Operator::Filter(FilterOperator { - predicate: ScalarExpression::Constant(DataValue::Boolean(true)), + predicate: arena + .alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))), is_optimized: false, having: false, }), diff --git a/src/optimizer/rule/normalization/min_max_top_k.rs b/src/optimizer/rule/normalization/min_max_top_k.rs index b1ce260b..47a174c9 100644 --- a/src/optimizer/rule/normalization/min_max_top_k.rs +++ b/src/optimizer/rule/normalization/min_max_top_k.rs @@ -28,7 +28,7 @@ impl NormalizationRule for MinMaxToTopK { fn apply( &self, plan: &mut LogicalPlan, - _: &mut crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result { let Operator::Aggregate(op) = &plan.operator else { return Ok(false); @@ -37,7 +37,7 @@ impl NormalizationRule for MinMaxToTopK { return Ok(false); } - let ScalarExpression::AggCall { kind, args, .. } = &op.agg_calls[0] else { + let ScalarExpression::AggCall { kind, args, .. } = arena.expression(op.agg_calls[0]) else { return Ok(false); }; if args.len() != 1 { @@ -50,7 +50,7 @@ impl NormalizationRule for MinMaxToTopK { _ => return Ok(false), }; - let sort_field = SortField::new(args[0].clone(), asc, false); + let sort_field = SortField::new(args[0], asc, false); let already_topk = match only_child(plan) { Some(child) => match &child.operator { Operator::TopK(topk) => { @@ -140,7 +140,7 @@ mod tests { assert_eq!(topk.sort_fields.len(), 1); assert!(topk.sort_fields[0].asc); assert!(!topk.sort_fields[0].nulls_first); - let args = match &op.agg_calls[0] { + let args = match arena.expression(op.agg_calls[0]) { crate::expression::ScalarExpression::AggCall { args, .. } => args, _ => unreachable!("Aggregate should use AggCall"), }; diff --git a/src/optimizer/rule/normalization/mod.rs b/src/optimizer/rule/normalization/mod.rs index 08e8e820..cec699b1 100644 --- a/src/optimizer/rule/normalization/mod.rs +++ b/src/optimizer/rule/normalization/mod.rs @@ -13,7 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; -use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut}; +use crate::expression::visitor_mut::ExprVisitorMut; use crate::expression::{AliasType, ScalarExpression}; use crate::optimizer::core::rule::NormalizationRule; use crate::optimizer::rule::normalization::column_pruning::ColumnPruning; @@ -33,7 +33,8 @@ use crate::optimizer::rule::normalization::pushdown_predicates::{ use crate::optimizer::rule::normalization::simplification::ConstantCalculation; use crate::optimizer::rule::normalization::simplification::SimplifyFilter; use crate::optimizer::rule::normalization::top_k::TopK; -use crate::planner::LogicalPlan; +use crate::planner::{ExprRef, LogicalPlan}; +use std::collections::HashSet; mod column_pruning; mod combine_operators; mod compilation_in_advance; @@ -220,16 +221,16 @@ impl NormalizationRule for NormalizationRuleImpl { } } -pub(crate) fn strip_alias(expr: &ScalarExpression) -> &ScalarExpression { - match expr { +pub(crate) fn strip_alias(expr: ExprRef, arena: &crate::planner::PlanArena<'_>) -> ExprRef { + match arena.expression(expr) { ScalarExpression::Alias { expr, alias: AliasType::Name(_), - } => strip_alias(expr), + } => strip_alias(*expr, arena), ScalarExpression::Alias { alias: AliasType::Expr(alias_expr), .. - } => strip_alias(alias_expr), + } => strip_alias(*alias_expr, arena), _ => expr, } } @@ -248,40 +249,74 @@ pub(crate) fn remap_position(position: &mut usize, removed_positions: &[usize]) } } -struct PositionRemapper<'a> { - removed_positions: &'a [usize], +struct PositionRemapper<'positions, 'visited> { + removed_positions: &'positions [usize], + visited: &'visited mut HashSet, } -impl<'a> ExprVisitorMut<'a> for PositionRemapper<'_> { - fn visit(&mut self, expr: &'a mut ScalarExpression) -> Result<(), DatabaseError> { - match expr { - ScalarExpression::ColumnRef { position, .. } => { - remap_position(position, self.removed_positions); - Ok(()) - } - ScalarExpression::Alias { expr, alias } => match alias { - AliasType::Expr(alias_expr) => self.visit(alias_expr), - AliasType::Name(_) => self.visit(expr), - }, - _ => walk_mut_expr(self, expr), +impl<'positions, 'visited> PositionRemapper<'positions, 'visited> { + pub(super) fn new( + removed_positions: &'positions [usize], + visited: &'visited mut HashSet, + ) -> Self { + visited.clear(); + Self { + removed_positions, + visited, + } + } +} + +impl ExprVisitorMut for PositionRemapper<'_, '_> { + fn visit_expression_ref( + &mut self, + expr: &mut ExprRef, + _arena: &mut crate::planner::PlanArena<'_>, + ) -> Result { + Ok(self.visited.insert(*expr)) + } + + fn visit_column_ref( + &mut self, + _column: &mut crate::catalog::ColumnRef, + position: &mut usize, + _arena: &mut crate::planner::PlanArena<'_>, + ) -> Result<(), DatabaseError> { + remap_position(position, self.removed_positions); + Ok(()) + } + + fn visit_alias( + &mut self, + expr: &mut ExprRef, + alias: &mut AliasType, + arena: &mut crate::planner::PlanArena<'_>, + ) -> Result<(), DatabaseError> { + match alias { + AliasType::Expr(alias_expr) => self.visit(alias_expr, arena), + AliasType::Name(_) => self.visit(expr, arena), } } } pub(crate) fn remap_expr_positions( - expr: &mut ScalarExpression, + mut expr: ExprRef, removed_positions: &[usize], + visited: &mut HashSet, + arena: &mut crate::planner::PlanArena<'_>, ) -> Result<(), DatabaseError> { - PositionRemapper { removed_positions }.visit(expr) + PositionRemapper::new(removed_positions, visited).visit(&mut expr, arena) } pub(crate) fn remap_exprs_positions<'a>( - exprs: impl IntoIterator, + exprs: impl IntoIterator, removed_positions: &[usize], + visited: &mut HashSet, + arena: &mut crate::planner::PlanArena<'_>, ) -> Result<(), DatabaseError> { - let mut remapper = PositionRemapper { removed_positions }; + let mut remapper = PositionRemapper::new(removed_positions, visited); for expr in exprs { - remapper.visit(expr)?; + remapper.visit(expr, arena)?; } Ok(()) } diff --git a/src/optimizer/rule/normalization/parameterized_index.rs b/src/optimizer/rule/normalization/parameterized_index.rs index cffb6390..5bd00392 100644 --- a/src/optimizer/rule/normalization/parameterized_index.rs +++ b/src/optimizer/rule/normalization/parameterized_index.rs @@ -19,7 +19,7 @@ use crate::optimizer::core::rule::NormalizationRule; use crate::planner::operator::mark_apply::{MarkApplyKind, MarkApplyQuantifier}; use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::{Operator, PhysicalOption, PlanImpl}; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::types::index::{IndexLookup, IndexType}; use crate::types::tuple::Schema; @@ -48,7 +48,7 @@ impl NormalizationRule for ParameterizeMarkApply { _ => return Ok(false), }; - let changed = op.parameterized_probe().cloned() != new_probe; + let changed = op.parameterized_probe().copied() != new_probe; op.set_parameterized_probe(new_probe); Ok(changed) } @@ -56,16 +56,16 @@ impl NormalizationRule for ParameterizeMarkApply { fn find_parameterized_probe( kind: MarkApplyKind, - predicates: &[ScalarExpression], + predicates: &[ExprRef], left_schema: &Schema, right_schema: &Schema, arena: &crate::planner::PlanArena, -) -> Result, DatabaseError> { +) -> Result, DatabaseError> { match kind { MarkApplyKind::Exists => { for predicate in predicates { if let Some(probe) = - extract_parameterized_probe(predicate, left_schema, right_schema, arena)? + extract_parameterized_probe(*predicate, left_schema, right_schema, arena)? { return Ok(Some(probe)); } @@ -74,7 +74,7 @@ fn find_parameterized_probe( } MarkApplyKind::Quantified(MarkApplyQuantifier::Any) => { if let Some(predicate) = predicates.first() { - extract_parameterized_probe(predicate, left_schema, right_schema, arena) + extract_parameterized_probe(*predicate, left_schema, right_schema, arena) } else { Ok(None) } @@ -84,12 +84,12 @@ fn find_parameterized_probe( } fn extract_parameterized_probe( - predicate: &ScalarExpression, + predicate: ExprRef, left_schema: &Schema, right_schema: &Schema, arena: &crate::planner::PlanArena, -) -> Result, DatabaseError> { - match predicate.unpack_alias_ref() { +) -> Result, DatabaseError> { + match predicate.unpack_alias_ref(arena) { ScalarExpression::Binary { op: BinaryOperator::Eq, left_expr, @@ -97,8 +97,8 @@ fn extract_parameterized_probe( .. } => { if let Some(probe) = extract_parameterized_probe_side( - left_expr, - right_expr, + *left_expr, + *right_expr, left_schema, right_schema, arena, @@ -106,8 +106,8 @@ fn extract_parameterized_probe( return Ok(Some(probe)); } extract_parameterized_probe_side( - right_expr, - left_expr, + *right_expr, + *left_expr, left_schema, right_schema, arena, @@ -118,13 +118,16 @@ fn extract_parameterized_probe( } fn extract_parameterized_probe_side( - right_expr: &ScalarExpression, - left_expr: &ScalarExpression, + right_expr: ExprRef, + left_expr: ExprRef, left_schema: &Schema, right_schema: &Schema, arena: &crate::planner::PlanArena, -) -> Result, DatabaseError> { - let Some((right_column, _)) = right_expr.unpack_alias_ref().unpack_bound_col(false) else { +) -> Result, DatabaseError> { + let Some((right_column, _)) = right_expr + .unpack_alias(arena) + .unpack_bound_col(arena, false) + else { return Ok(None); }; @@ -142,7 +145,7 @@ fn extract_parameterized_probe_side( return Ok(None); } - Ok(Some((right_column, left_expr.clone()))) + Ok(Some((right_column, left_expr))) } fn parameterize_right_subtree( @@ -253,14 +256,18 @@ mod tests { )) } - fn eq(left: ScalarExpression, right: ScalarExpression) -> ScalarExpression { - ScalarExpression::Binary { + fn expr(arena: &mut PlanArena, column: ColumnRef) -> ExprRef { + arena.alloc_expression(ScalarExpression::column_expr(column, 0)) + } + + fn eq(arena: &mut PlanArena, left: ExprRef, right: ExprRef) -> ExprRef { + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Eq, - left_expr: Box::new(left), - right_expr: Box::new(right), + left_expr: left, + right_expr: right, evaluator: None, ty: LogicalType::Boolean, - } + }) } #[test] @@ -274,6 +281,9 @@ mod tests { let right_schema = vec![right]; let overlapping_schema = vec![left, right]; + let all_right = expr(&mut arena, right); + let all_left = expr(&mut arena, left); + let all_predicate = eq(&mut arena, all_right, all_left); assert!(find_parameterized_probe( MarkApplyKind::Quantified(MarkApplyQuantifier::Any), &[], @@ -284,22 +294,19 @@ mod tests { .is_none()); assert!(find_parameterized_probe( MarkApplyKind::Quantified(MarkApplyQuantifier::All), - &[eq( - ScalarExpression::column_expr(right, 0), - ScalarExpression::column_expr(left, 0), - )], + &[all_predicate], &left_schema, &right_schema, &arena, )? .is_none()); + let exists_right = expr(&mut arena, right); + let exists_left = expr(&mut arena, left); + let exists_predicate = eq(&mut arena, exists_right, exists_left); let predicates = vec![ - ScalarExpression::from(true), - eq( - ScalarExpression::column_expr(right, 0), - ScalarExpression::column_expr(left, 0), - ), + arena.alloc_expression(ScalarExpression::from(true)), + exists_predicate, ]; let probe = find_parameterized_probe( MarkApplyKind::Exists, @@ -311,28 +318,28 @@ mod tests { .expect("right = left should be parameterizable"); assert_eq!(probe.0, right); + let outside_right = expr(&mut arena, right); + let outside_expr = expr(&mut arena, outside); + let outside_predicate = eq(&mut arena, outside_right, outside_expr); assert!(extract_parameterized_probe( - &eq( - ScalarExpression::column_expr(right, 0), - ScalarExpression::column_expr(outside, 0), - ), + outside_predicate, &left_schema, &right_schema, &arena, )? .is_none()); + let overlap_left = expr(&mut arena, right); + let overlap_right = expr(&mut arena, right); + let overlap_predicate = eq(&mut arena, overlap_left, overlap_right); assert!(extract_parameterized_probe( - &eq( - ScalarExpression::column_expr(right, 0), - ScalarExpression::column_expr(right, 0), - ), + overlap_predicate, &overlapping_schema, &right_schema, &arena, )? .is_none()); assert!(extract_parameterized_probe( - &ScalarExpression::from(false), + arena.alloc_expression(ScalarExpression::from(false)), &left_schema, &right_schema, &arena, @@ -355,7 +362,7 @@ mod tests { let mut filter = LogicalPlan::new( Operator::Filter(FilterOperator { - predicate: ScalarExpression::from(true), + predicate: arena.alloc_expression(ScalarExpression::from(true)), is_optimized: false, having: false, }), diff --git a/src/optimizer/rule/normalization/pushdown_predicates.rs b/src/optimizer/rule/normalization/pushdown_predicates.rs index 5494cc52..104ad37c 100644 --- a/src/optimizer/rule/normalization/pushdown_predicates.rs +++ b/src/optimizer/rule/normalization/pushdown_predicates.rs @@ -21,41 +21,38 @@ use crate::optimizer::plan_utils::{replace_with_only_child, wrap_child_with}; use crate::planner::operator::filter::FilterOperator; use crate::planner::operator::join::{JoinCondition, JoinType}; use crate::planner::operator::{Operator, SortOption}; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::index::{IndexInfo, IndexLookup, IndexMetaRef, IndexType}; use crate::types::value::DataValue; use crate::types::LogicalType; use std::ops::Bound; -use std::{borrow::Cow, mem, slice}; +use std::{mem, slice}; const EMPTY_SCHEMA: [crate::catalog::ColumnRef; 0] = []; -type ClassifiedJoinFilters = ( - Option, - Option, - Option, - usize, -); +type ClassifiedJoinFilters = (Option, Option, Option, usize); fn split_conjunctive_predicates( - expr: &ScalarExpression, - f: &mut impl FnMut(ScalarExpression) -> Result<(), DatabaseError>, + expr: ExprRef, + arena: &mut PlanArena<'_>, + f: &mut impl FnMut(ExprRef, &mut PlanArena<'_>) -> Result<(), DatabaseError>, ) -> Result<(), DatabaseError> { - match expr { + match arena.expression(expr) { ScalarExpression::Binary { op: BinaryOperator::And, left_expr, right_expr, .. } => { - split_conjunctive_predicates(left_expr, f)?; - split_conjunctive_predicates(right_expr, f) + let (left_expr, right_expr) = (*left_expr, *right_expr); + split_conjunctive_predicates(left_expr, arena, f)?; + split_conjunctive_predicates(right_expr, arena, f) } - _ => f(expr.clone()), + _ => f(expr, arena), } } fn classify_join_filters( - predicate: &ScalarExpression, + predicate: ExprRef, childrens: &mut Childrens, arena: &mut crate::planner::PlanArena, ) -> Result { @@ -68,28 +65,28 @@ fn classify_join_filters( Childrens::None => (&EMPTY_SCHEMA, &EMPTY_SCHEMA), }; let left_len = left_columns.len(); - let append_filter = |slot: &mut Option, expr| { + let append_filter = |slot: &mut Option, expr, arena: &mut PlanArena<'_>| { *slot = Some(match slot.take() { - Some(current) => ScalarExpression::Binary { + Some(current) => arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(current), - right_expr: Box::new(expr), + left_expr: current, + right_expr: expr, evaluator: None, ty: LogicalType::Boolean, - }, + }), None => expr, }); }; let mut left_filter = None; let mut right_filter = None; let mut common_filter = None; - split_conjunctive_predicates(predicate, &mut |expr| { + split_conjunctive_predicates(predicate, arena, &mut |expr, arena| { if expr.all_referenced_columns(arena, |_, column| left_columns.contains(column))? { - append_filter(&mut left_filter, expr); + append_filter(&mut left_filter, expr, arena); } else if expr.all_referenced_columns(arena, |_, column| right_columns.contains(column))? { - append_filter(&mut right_filter, expr); + append_filter(&mut right_filter, expr, arena); } else { - append_filter(&mut common_filter, expr); + append_filter(&mut common_filter, expr, arena); } Ok(()) })?; @@ -99,17 +96,20 @@ fn classify_join_filters( /// reduce filters into a filter, and then build a new LogicalFilter node with input child. /// if filters is empty, return the input child. fn reduce_filters( - filters: impl IntoIterator, + filters: impl IntoIterator, having: bool, + arena: &mut PlanArena<'_>, ) -> Option { filters .into_iter() - .reduce(|a, b| ScalarExpression::Binary { - op: BinaryOperator::And, - left_expr: Box::new(a), - right_expr: Box::new(b), - evaluator: None, - ty: LogicalType::Boolean, + .reduce(|a, b| { + arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::And, + left_expr: a, + right_expr: b, + evaluator: None, + ty: LogicalType::Boolean, + }) }) .map(|f| FilterOperator { predicate: f, @@ -165,12 +165,14 @@ impl NormalizationRule for PushPredicateThroughJoin { } let (left_filter, mut right_filter, common_filter, left_len) = - classify_join_filters(&filter_op.predicate, join_plan.childrens.as_mut(), arena)?; + classify_join_filters(filter_op.predicate, join_plan.childrens.as_mut(), arena)?; let mut new_ops = (None, None, None); match join_type { JoinType::Inner => { - if let Some(left_filter_op) = reduce_filters(left_filter, filter_op.having) { + if let Some(left_filter_op) = + reduce_filters(left_filter, filter_op.having, arena) + { new_ops.0 = Some(Operator::Filter(left_filter_op)); } @@ -178,21 +180,26 @@ impl NormalizationRule for PushPredicateThroughJoin { PositionShift { delta: -(left_len as isize), } - .visit(expr)?; + .visit(expr, arena)?; } - if let Some(right_filter_op) = reduce_filters(right_filter, filter_op.having) { + if let Some(right_filter_op) = + reduce_filters(right_filter, filter_op.having, arena) + { new_ops.1 = Some(Operator::Filter(right_filter_op)); } - new_ops.2 = - reduce_filters(common_filter, filter_op.having).map(Operator::Filter); + new_ops.2 = reduce_filters(common_filter, filter_op.having, arena) + .map(Operator::Filter); } JoinType::LeftOuter => { - if let Some(left_filter_op) = reduce_filters(left_filter, filter_op.having) { + if let Some(left_filter_op) = + reduce_filters(left_filter, filter_op.having, arena) + { new_ops.0 = Some(Operator::Filter(left_filter_op)); } new_ops.2 = reduce_filters( common_filter.into_iter().chain(right_filter), filter_op.having, + arena, ) .map(Operator::Filter); } @@ -201,14 +208,17 @@ impl NormalizationRule for PushPredicateThroughJoin { PositionShift { delta: -(left_len as isize), } - .visit(expr)?; + .visit(expr, arena)?; } - if let Some(right_filter_op) = reduce_filters(right_filter, filter_op.having) { + if let Some(right_filter_op) = + reduce_filters(right_filter, filter_op.having, arena) + { new_ops.1 = Some(Operator::Filter(right_filter_op)); } new_ops.2 = reduce_filters( common_filter.into_iter().chain(left_filter), filter_op.having, + arena, ) .map(Operator::Filter); } @@ -287,16 +297,13 @@ impl NormalizationRule for PushPredicateIntoScan { else { return Err(DatabaseError::InvalidIndex); }; - let index_meta = arena.index(*meta); - let detached = match index_meta.ty { + let index_type = arena.index(*meta).ty; + let detached = match index_type { IndexType::PrimaryKey { is_multiple: false } | IndexType::Unique - | IndexType::Normal => RangeDetacher::new( - index_meta.table_name.as_ref(), - &index_meta.column_ids[0], - arena, - ) - .detach(&filter_op.predicate)?, + | IndexType::Normal => { + RangeDetacher::for_index(*meta, 0, arena).detach(filter_op.predicate)? + } IndexType::PrimaryKey { is_multiple: true } | IndexType::Composite => { Self::composite_range(filter_op, *meta, ignore_prefix_len, arena)? } @@ -358,25 +365,25 @@ impl PushPredicateIntoScan { op: &FilterOperator, meta: IndexMetaRef, ignore_prefix_len: &mut usize, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result, DatabaseError> { - let meta = arena.index(meta); + let column_count = arena.index(meta).column_ids.len(); let mut res = None; - let mut eq_ranges = Vec::with_capacity(meta.column_ids.len()); + let mut eq_ranges = Vec::with_capacity(column_count); let mut apply_column_count = 0; - let mut residual = Some(Cow::Borrowed(&op.predicate)); + let mut residual = Some(op.predicate); - for column_id in meta.column_ids.iter() { + for position in 0..column_count { let Some(predicate) = residual.take() else { break; }; - let Some(detached) = RangeDetacher::new(meta.table_name.as_ref(), column_id, arena) - .detach(predicate.as_ref())? + let Some(detached) = + RangeDetacher::for_index(meta, position, arena).detach(predicate)? else { residual = Some(predicate); break; }; - residual = detached.residual.map(Cow::Owned); + residual = detached.residual; let range = detached.range; apply_column_count += 1; @@ -396,7 +403,7 @@ impl PushPredicateIntoScan { } } let range = res.map(|range| { - if range.only_eq() && apply_column_count != meta.column_ids.len() { + if range.only_eq() && apply_column_count != column_count { fn eq_to_scope(range: Range) -> Range { match range { Range::Eq(DataValue::Tuple(values, _)) => { @@ -419,10 +426,7 @@ impl PushPredicateIntoScan { } range }); - Ok(range.map(|range| DetachedPredicate { - range, - residual: residual.map(|predicate| predicate.into_owned()), - })) + Ok(range.map(|range| DetachedPredicate { range, residual })) } } @@ -454,7 +458,7 @@ impl NormalizationRule for PushJoinPredicateIntoScan { }; let (left_filter, mut right_filter, common_filter, left_len) = - classify_join_filters(&filter_expr, plan.childrens.as_mut(), arena)?; + classify_join_filters(filter_expr, plan.childrens.as_mut(), arena)?; let (push_left, push_right) = match join_type { JoinType::Inner => (true, true), @@ -465,7 +469,7 @@ impl NormalizationRule for PushJoinPredicateIntoScan { let mut new_ops = (None, None); let left_remain = if push_left { - if let Some(filter_op) = reduce_filters(left_filter, false) { + if let Some(filter_op) = reduce_filters(left_filter, false, arena) { new_ops.0 = Some(Operator::Filter(filter_op)); } None @@ -478,9 +482,9 @@ impl NormalizationRule for PushJoinPredicateIntoScan { PositionShift { delta: -(left_len as isize), } - .visit(expr)?; + .visit(expr, arena)?; } - if let Some(filter_op) = reduce_filters(right_filter, false) { + if let Some(filter_op) = reduce_filters(right_filter, false, arena) { new_ops.1 = Some(Operator::Filter(filter_op)); } None @@ -502,6 +506,7 @@ impl NormalizationRule for PushJoinPredicateIntoScan { .chain(left_remain) .chain(right_remain), false, + arena, ) .map(|op| op.predicate); let filter_changed = match &join_filter { @@ -543,7 +548,7 @@ mod tests { use crate::planner::operator::join::{JoinCondition, JoinType}; use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::{Operator, SortOption}; - use crate::planner::{Childrens, LogicalPlan, PlanArena}; + use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::index::{IndexInfo, IndexLookup, IndexMeta, IndexType}; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -573,28 +578,32 @@ mod tests { } fn cmp_predicate( + arena: &mut PlanArena, op: BinaryOperator, column: crate::catalog::ColumnRef, position: usize, value: i32, - ) -> ScalarExpression { - ScalarExpression::Binary { + ) -> ExprRef { + let left_expr = arena.alloc_expression(ScalarExpression::column_expr(column, position)); + let right_expr = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(value))); + arena.alloc_expression(ScalarExpression::Binary { op, - left_expr: Box::new(ScalarExpression::column_expr(column, position)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(value))), + left_expr, + right_expr, evaluator: None, ty: LogicalType::Boolean, - } + }) } - fn and_predicate(left: ScalarExpression, right: ScalarExpression) -> ScalarExpression { - ScalarExpression::Binary { + fn and_predicate(arena: &mut PlanArena, left: ExprRef, right: ExprRef) -> ExprRef { + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, - left_expr: Box::new(left), - right_expr: Box::new(right), + left_expr: left, + right_expr: right, evaluator: None, ty: LogicalType::Boolean, - } + }) } #[test] @@ -661,16 +670,13 @@ mod tests { name: "idx_c1_c2_c3".to_string(), ty: IndexType::Composite, }); - let predicate = and_predicate( - and_predicate( - cmp_predicate(BinaryOperator::Eq, c1, 0, 1), - cmp_predicate(BinaryOperator::Eq, c2, 1, 2), - ), - and_predicate( - cmp_predicate(BinaryOperator::Eq, c3, 2, 3), - cmp_predicate(BinaryOperator::Eq, c4, 3, 4), - ), - ); + let c1_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c1, 0, 1); + let c2_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c2, 1, 2); + let c3_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c3, 2, 3); + let c4_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c4, 3, 4); + let left = and_predicate(&mut arena, c1_eq, c2_eq); + let right = and_predicate(&mut arena, c3_eq, c4_eq); + let predicate = and_predicate(&mut arena, left, right); let filter = FilterOperator { predicate, is_optimized: false, @@ -682,7 +688,7 @@ mod tests { &filter, index_meta, &mut ignore_prefix_len, - &arena, + &mut arena, )? .expect("composite prefix should be consumed"); @@ -699,8 +705,8 @@ mod tests { )) ); let residual = detached.residual.expect("c4 predicate should remain"); - let residual_detached = RangeDetacher::new(table_name.as_ref(), &4, &arena) - .detach(&residual)? + let residual_detached = RangeDetacher::new(table_name.as_ref(), &4, &mut arena) + .detach(residual)? .expect("residual should be the c4 predicate"); assert_eq!(residual_detached.range, Range::Eq(DataValue::Int32(4))); assert_eq!(residual_detached.residual, None); @@ -730,13 +736,11 @@ mod tests { name: "idx_c1_c2_c3".to_string(), ty: IndexType::Composite, }); - let predicate = and_predicate( - and_predicate( - cmp_predicate(BinaryOperator::Eq, c1, 0, 1), - cmp_predicate(BinaryOperator::Gt, c2, 1, 2), - ), - cmp_predicate(BinaryOperator::Eq, c3, 2, 3), - ); + let c1_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c1, 0, 1); + let c2_gt = cmp_predicate(&mut arena, BinaryOperator::Gt, c2, 1, 2); + let c3_eq = cmp_predicate(&mut arena, BinaryOperator::Eq, c3, 2, 3); + let prefix = and_predicate(&mut arena, c1_eq, c2_gt); + let predicate = and_predicate(&mut arena, prefix, c3_eq); let filter = FilterOperator { predicate, is_optimized: false, @@ -748,7 +752,7 @@ mod tests { &filter, index_meta, &mut ignore_prefix_len, - &arena, + &mut arena, )? .expect("composite prefix should be consumed"); @@ -764,8 +768,8 @@ mod tests { } ); let residual = detached.residual.expect("c3 predicate should remain"); - let residual_detached = RangeDetacher::new(table_name.as_ref(), &3, &arena) - .detach(&residual)? + let residual_detached = RangeDetacher::new(table_name.as_ref(), &3, &mut arena) + .detach(residual)? .expect("residual should be the c3 predicate"); assert_eq!(residual_detached.range, Range::Eq(DataValue::Int32(3))); assert_eq!(residual_detached.residual, None); @@ -832,7 +836,7 @@ mod tests { let scan_plan = LogicalPlan::new( Operator::TableScan(TableScanOperator { - table_name: table_name.clone(), + table_name, columns, limit: (None, None), index_infos: vec![ @@ -868,27 +872,9 @@ mod tests { Childrens::None, ); - let c1_gt = ScalarExpression::Binary { - op: BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::column_expr(c1_ref, 0)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(0))), - evaluator: None, - ty: LogicalType::Boolean, - }; - let c2_gt = ScalarExpression::Binary { - op: BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::column_expr(c2_ref, 1)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(0))), - evaluator: None, - ty: LogicalType::Boolean, - }; - let predicate = ScalarExpression::Binary { - op: BinaryOperator::And, - left_expr: Box::new(c1_gt), - right_expr: Box::new(c2_gt), - evaluator: None, - ty: LogicalType::Boolean, - }; + let c1_gt = cmp_predicate(&mut arena, BinaryOperator::Gt, c1_ref, 0, 0); + let c2_gt = cmp_predicate(&mut arena, BinaryOperator::Gt, c2_ref, 1, 0); + let predicate = and_predicate(&mut arena, c1_gt, c2_gt); let filter_plan = LogicalPlan::new( Operator::Filter(FilterOperator { @@ -974,7 +960,7 @@ mod tests { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(op) = &filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Lt, ty: LogicalType::Boolean, @@ -988,7 +974,7 @@ mod tests { let filter_op = filter_op.childrens.pop_only().childrens.pop_twins().0; if let Operator::Filter(op) = &filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Gt, ty: LogicalType::Boolean, @@ -1024,7 +1010,7 @@ mod tests { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(op) = &filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Gt, ty: LogicalType::Boolean, @@ -1038,7 +1024,7 @@ mod tests { let filter_op = filter_op.childrens.pop_only().childrens.pop_twins().1; if let Operator::Filter(op) = &filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Lt, ty: LogicalType::Boolean, @@ -1080,7 +1066,7 @@ mod tests { let (left_filter_op, right_filter_op) = join_op.childrens.pop_twins(); if let Operator::Filter(op) = &left_filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Gt, ty: LogicalType::Boolean, @@ -1093,7 +1079,7 @@ mod tests { } if let Operator::Filter(op) = &right_filter_op.operator { - match op.predicate { + match arena.expression(op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Lt, ty: LogicalType::Boolean, @@ -1148,7 +1134,7 @@ mod tests { let (left_child, right_child) = join_plan.childrens.pop_twins(); if let Operator::Filter(left_filter) = &left_child.operator { - match left_filter.predicate { + match arena.expression(left_filter.predicate) { ScalarExpression::Binary { op: BinaryOperator::Gt, ty: LogicalType::Boolean, @@ -1165,7 +1151,7 @@ mod tests { } if let Operator::Filter(right_filter) = &right_child.operator { - match right_filter.predicate { + match arena.expression(right_filter.predicate) { ScalarExpression::Binary { op: BinaryOperator::Lt, ty: LogicalType::Boolean, @@ -1276,7 +1262,7 @@ mod tests { Operator::Filter(ref op) => op, _ => unreachable!("right child should be a filter"), }; - match filter_op.predicate { + match arena.expression(filter_op.predicate) { ScalarExpression::Binary { op: BinaryOperator::Lt, ty: LogicalType::Boolean, diff --git a/src/optimizer/rule/normalization/simplification.rs b/src/optimizer/rule/normalization/simplification.rs index 308b1409..604d37ac 100644 --- a/src/optimizer/rule/normalization/simplification.rs +++ b/src/optimizer/rule/normalization/simplification.rs @@ -25,16 +25,16 @@ pub struct ConstantCalculation; pub(crate) fn constant_calculation_current( plan: &mut LogicalPlan, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { let mut calculator = ConstantCalculator::new(arena); - OperatorExprVisitorMut::new(&mut calculator).visit_operator(&mut plan.operator) + OperatorExprVisitorMut::new(&mut calculator, arena).visit_operator(&mut plan.operator) } impl ConstantCalculation { fn _apply( plan: &mut LogicalPlan, - arena: &crate::planner::PlanArena, + arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { constant_calculation_current(plan, arena)?; match plan.childrens.as_mut() { @@ -93,8 +93,8 @@ impl NormalizationRule for SimplifyFilter { return Ok(false); } } - ConstantCalculator::new(arena).visit(&mut filter_op.predicate)?; - Simplify::default().visit(&mut filter_op.predicate)?; + ConstantCalculator::new(arena).visit(&mut filter_op.predicate, arena)?; + Simplify::default().visit(&mut filter_op.predicate, arena)?; filter_op.is_optimized = true; return Ok(true); } @@ -157,19 +157,21 @@ mod test { .find_best(None, &mut arena)?; if let Operator::Project(project_op) = best_plan.clone().operator { let constant_expr = ScalarExpression::Constant(DataValue::Int32(3)); - if let ScalarExpression::Binary { right_expr, .. } = &project_op.exprs[0] { - assert_eq!(right_expr.as_ref(), &constant_expr); + if let ScalarExpression::Binary { right_expr, .. } = + arena.expression(project_op.exprs[0]) + { + assert_eq!(arena.expression(*right_expr), &constant_expr); } else { unreachable!(); } - assert_eq!(&project_op.exprs[1], &constant_expr); + assert_eq!(arena.expression(project_op.exprs[1]), &constant_expr); } else { unreachable!(); } let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(filter_op) = filter_op.operator { - let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &arena) - .detach(&filter_op.predicate)? + let range = RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut arena) + .detach(filter_op.predicate)? .map(|detached| detached.range) .unwrap(); assert_eq!( @@ -205,12 +207,12 @@ mod test { if let Operator::Project(project_op) = best_plan.operator { assert_eq!( - project_op.exprs[0], - ScalarExpression::Constant(DataValue::Int32(1)) + arena.expression(project_op.exprs[0]), + &ScalarExpression::Constant(DataValue::Int32(1)) ); assert_eq!( - project_op.exprs[1], - ScalarExpression::Constant(DataValue::Int32(3)) + arena.expression(project_op.exprs[1]), + &ScalarExpression::Constant(DataValue::Int32(3)) ); } else { unreachable!(); @@ -219,8 +221,8 @@ mod test { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(filter_op) = filter_op.operator { assert_eq!( - filter_op.predicate, - ScalarExpression::Constant(DataValue::Boolean(true)) + arena.expression(filter_op.predicate), + &ScalarExpression::Constant(DataValue::Boolean(true)) ); } else { unreachable!(); @@ -229,6 +231,40 @@ mod test { Ok(()) } + #[test] + fn test_constant_is_null_elimination() -> Result<(), DatabaseError> { + use crate::expression::simplify::Simplify; + use crate::expression::visitor_mut::ExprVisitorMut; + use crate::planner::TableArenaCell; + + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + for (value, negated, expected) in [ + (DataValue::Null, false, true), + (DataValue::Null, true, false), + (DataValue::Int32(1), false, false), + (DataValue::Int32(1), true, true), + ] { + let value = arena.alloc_expression(ScalarExpression::Constant(value)); + let mut expression = arena.alloc_expression(ScalarExpression::IsNull { + negated, + expr: value, + }); + + assert_eq!( + expression.unpack_val(&arena), + Some(DataValue::Boolean(expected)) + ); + Simplify::default().visit(&mut expression, &mut arena)?; + assert_eq!( + arena.expression(expression), + &ScalarExpression::Constant(DataValue::Boolean(expected)) + ); + } + + Ok(()) + } + #[test] fn test_simplify_filter_single_column() -> Result<(), DatabaseError> { let table_state = build_t1_table()?; @@ -275,8 +311,8 @@ mod test { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(filter_op) = filter_op.operator { Ok( - RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &arena) - .detach(&filter_op.predicate)? + RangeDetacher::new("t1", table_state.column_id_by_name("c1"), &mut arena) + .detach(filter_op.predicate)? .map(|detached| detached.range), ) } else { @@ -350,27 +386,32 @@ mod test { let c2_ref = table_state.table.get_column_by_name("c2").unwrap(); // -(c1 + 1) > c2 => c1 < -c2 - 1 - assert_eq!( - filter_op.predicate, - ScalarExpression::Binary { + let c1 = arena.alloc_expression(ScalarExpression::column_expr(c1_ref, 0)); + let one = arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(1))); + let plus = arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::Plus, + left_expr: c1, + right_expr: one, + evaluator: None, + ty: LogicalType::Integer, + }); + let minus = arena.alloc_expression(ScalarExpression::Unary { + op: UnaryOperator::Minus, + expr: plus, + evaluator: None, + ty: LogicalType::Integer, + }); + let c2 = arena.alloc_expression(ScalarExpression::column_expr(c2_ref, 1)); + assert!(filter_op.predicate.eq_ignore_colref_pos( + arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::Gt, - left_expr: Box::new(ScalarExpression::Unary { - op: UnaryOperator::Minus, - expr: Box::new(ScalarExpression::Binary { - op: BinaryOperator::Plus, - left_expr: Box::new(ScalarExpression::column_expr(c1_ref, 0)), - right_expr: Box::new(ScalarExpression::Constant(DataValue::Int32(1))), - evaluator: None, - ty: LogicalType::Integer, - }), - evaluator: None, - ty: LogicalType::Integer, - }), - right_expr: Box::new(ScalarExpression::column_expr(c2_ref, 1)), + left_expr: minus, + right_expr: c2, evaluator: None, ty: LogicalType::Boolean, - } - ) + }), + &arena, + )); } else { unreachable!() } @@ -394,7 +435,7 @@ mod test { let filter_op = best_plan.childrens.pop_only(); if let Operator::Filter(filter_op) = filter_op.operator { Ok(RangeDetacher::new("t1", column_id, arena) - .detach(&filter_op.predicate)? + .detach(filter_op.predicate)? .map(|detached| detached.range)) } else { Ok(None) diff --git a/src/orm/ddl.rs b/src/orm/ddl.rs index 5cafff4e..e213e7b3 100644 --- a/src/orm/ddl.rs +++ b/src/orm/ddl.rs @@ -140,7 +140,7 @@ impl Database { /// when the underlying DDL supports them. Primary-key changes and unique /// constraint changes still return an error so you can handle them manually. pub fn migrate(&mut self) -> Result<(), DatabaseError> { - let columns = M::columns(); + let columns = M::columns(self.state.table_arena().borrow_mut()); if columns.is_empty() { return Err(DatabaseError::UnsupportedStmt( "ORM migration requires Model::columns(); #[derive(Model)] provides it automatically" @@ -173,7 +173,11 @@ impl Database { .find(|column| column.desc().is_primary()) .ok_or(DatabaseError::PrimaryKeyNotFound)?; if table_primary_key.name() != model_primary_key.name() - || !model_column_matches_catalog(model_primary_key, &table_primary_key)? + || !model_column_matches_catalog( + model_primary_key, + &table_primary_key, + &PlanArena::new(self.state.table_arena()), + )? { return Err(DatabaseError::InvalidValue(::std::format!( "ORM migration does not support changing the primary key for table `{}`", @@ -187,7 +191,7 @@ impl Database { let mut handled_current = BTreeMap::new(); let mut handled_model = BTreeMap::new(); - for column in columns { + for column in &columns { let Some(current_column) = current_columns.get(column.name()) else { continue; }; @@ -207,7 +211,11 @@ impl Database { M::table_name(), ))); } - if model_column_matches_catalog(column, current_column)? { + if model_column_matches_catalog( + column, + current_column, + &PlanArena::new(self.state.table_arena()), + )? { continue; } @@ -223,14 +231,17 @@ impl Database { )?; } - if model_column_default(column)? != catalog_column_default(current_column)? { + let arena = PlanArena::new(self.state.table_arena()); + if model_column_default(column, &arena)? + != catalog_column_default(current_column, &arena)? + { execute_change_column( self, M::table_name(), column.name(), column.name(), column.datatype().clone(), - match column.desc().default.clone() { + match column.desc().default { Some(expr) => DefaultChange::Set(expr), None => DefaultChange::Drop, }, @@ -275,7 +286,11 @@ impl Database { .iter() .filter(|column| !column.desc().is_primary()) { - if model_column_rename_compatible(model_column, column)? { + if model_column_rename_compatible( + model_column, + column, + &PlanArena::new(self.state.table_arena()), + )? { candidates.push(column); } } @@ -288,7 +303,11 @@ impl Database { .iter() .filter(|other| !other.desc().is_primary()) { - if model_column_rename_compatible(other, current_column)? { + if model_column_rename_compatible( + other, + current_column, + &PlanArena::new(self.state.table_arena()), + )? { reverse_candidates.push(other); } } @@ -332,7 +351,7 @@ impl Database { execute_drop_column(self, M::table_name(), column.name())?; } - for column in columns { + for column in &columns { if handled_model.contains_key(column.name()) || current_columns.contains_key(column.name()) { @@ -409,7 +428,7 @@ fn execute_create_table( database: &mut Database, if_not_exists: bool, ) -> Result<(), DatabaseError> { - let columns = M::columns().to_vec(); + let columns = M::columns(database.state.table_arena().borrow_mut()); database.execute_mut("ORM CREATE TABLE", &[], move |binder, _| { binder.bind_create_table(M::table_name().into(), columns, if_not_exists) }) diff --git a/src/orm/mod.rs b/src/orm/mod.rs index d0df33f8..f1249f64 100644 --- a/src/orm/mod.rs +++ b/src/orm/mod.rs @@ -11,12 +11,12 @@ use crate::db::{ use crate::errors::DatabaseError; pub use crate::expression::agg::AggKind; use crate::expression::window::WindowFunctionKind; -use crate::expression::{self, AliasType, ScalarExpression}; +use crate::expression::{self, AliasType, ScalarExpression, TypeCast}; use crate::planner::operator::alter_table::change_column::{DefaultChange, NotNullChange}; use crate::planner::operator::join::JoinType; use crate::planner::operator::mark_apply::MarkApplyQuantifier; use crate::planner::operator::sort::SortField; -use crate::planner::{LogicalPlan, PlanArena}; +use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::{Storage, Transaction}; use crate::types::tuple::{SchemaView, Tuple}; use crate::types::value::DataValue; @@ -106,7 +106,7 @@ pub struct FieldSort { /// Partitioning and ordering for an ORM window expression. #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct WindowSpec { - partition_by: Vec, + partition_by: Vec, order_by: Vec, } @@ -204,7 +204,7 @@ where fn bind_scalar( self, scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, T, A>, - ) -> Result; + ) -> Result; } impl<'bind, 'parent, 'arena, T, A, M, V> BindOrmScalar<'bind, 'parent, 'arena, T, A> for Field @@ -215,7 +215,7 @@ where fn bind_scalar( self, scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, T, A>, - ) -> Result { + ) -> Result { scope.column(self).map(CtxExpression::into_scalar) } } @@ -228,8 +228,8 @@ where fn bind_scalar( self, _scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, T, A>, - ) -> Result { - Ok(self) + ) -> Result { + Ok(_scope.arena.alloc_expression(self)) } } @@ -242,7 +242,7 @@ where fn bind_scalar( self, _scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, T, A>, - ) -> Result { + ) -> Result { Ok(self.into_scalar()) } } @@ -296,9 +296,9 @@ where { fn bind_sort<'scope>( self, - _scope: &'scope mut ExprBindScope<'scope, 'bind, 'parent, 'arena, T, A>, + scope: &'scope mut ExprBindScope<'scope, 'bind, 'parent, 'arena, T, A>, ) -> Result { - Ok(self.into()) + Ok(SortField::from(scope.arena.alloc_expression(self))) } } @@ -330,8 +330,14 @@ where } #[doc(hidden)] -pub trait IntoOrmScalarExpression { - fn into_orm_scalar(self) -> ScalarExpression; +pub enum OrmExpression { + Bound(ExprRef), + Unbound(ScalarExpression), +} + +#[doc(hidden)] +pub trait IntoOrmExpression { + fn into_orm_expression(self) -> OrmExpression; } impl WindowSpec { @@ -339,8 +345,8 @@ impl WindowSpec { Self::default() } - pub fn partition_by(mut self, expr: impl IntoOrmScalarExpression) -> Self { - self.partition_by.push(expr.into_orm_scalar()); + pub fn partition_by(mut self, expr: impl Into) -> Self { + self.partition_by.push(expr.into()); self } @@ -350,12 +356,40 @@ impl WindowSpec { } } -impl IntoOrmScalarExpression for E +impl<'bind, 'parent, 'arena, T, A> From> for ExprRef +where + T: Transaction, + A: AsRef<[(&'static str, DataValue)]>, +{ + fn from(expr: CtxExpression<'bind, 'parent, 'arena, T, A>) -> Self { + expr.into_scalar() + } +} + +impl From for OrmExpression { + fn from(expr: ExprRef) -> Self { + Self::Bound(expr) + } +} + +impl From for OrmExpression { + fn from(expr: ScalarExpression) -> Self { + Self::Unbound(expr) + } +} + +impl IntoOrmExpression for ExprRef { + fn into_orm_expression(self) -> OrmExpression { + OrmExpression::Bound(self) + } +} + +impl IntoOrmExpression for E where E: Into, { - fn into_orm_scalar(self) -> ScalarExpression { - self.into() + fn into_orm_expression(self) -> OrmExpression { + OrmExpression::Unbound(self.into()) } } @@ -368,7 +402,7 @@ where fn bind_scalar_list( self, scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, T, A>, - ) -> Result, DatabaseError>; + ) -> Result, DatabaseError>; } macro_rules! impl_bind_orm_scalar_list { @@ -385,7 +419,7 @@ macro_rules! impl_bind_orm_scalar_list { fn bind_scalar_list( self, scope: &mut ExprBindScope<'_, 'bind, 'parent, 'arena, Tx, Args>, - ) -> Result, DatabaseError> { + ) -> Result, DatabaseError> { let ($($name,)+) = self; Ok(vec![ $($name.bind_scalar(scope)?,)+ @@ -467,8 +501,22 @@ where } } - fn wrap(self, expr: ScalarExpression) -> CtxExpression<'bind, 'parent, 'arena, T, A> { - CtxExpression { expr, scope: self } + fn wrap(self, expr: impl Into) -> CtxExpression<'bind, 'parent, 'arena, T, A> { + CtxExpression { + expr: self.bind(expr), + scope: self, + } + } + + fn alloc(self, expr: ScalarExpression) -> CtxExpression<'bind, 'parent, 'arena, T, A> { + self.wrap(expr) + } + + fn bind(self, expr: impl Into) -> ExprRef { + match expr.into() { + OrmExpression::Bound(expr) => expr, + OrmExpression::Unbound(expr) => self.arena().alloc_expression(expr), + } } #[allow(clippy::mut_from_ref)] @@ -490,54 +538,54 @@ where fn binary( self, - left: ScalarExpression, + left: ExprRef, op: expression::BinaryOperator, - right: ScalarExpression, + right: ExprRef, ) -> Result, DatabaseError> { self.binder() .bind_binary_op_expr(left, right, op, self.arena()) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn unary( self, op: expression::UnaryOperator, - expr: ScalarExpression, + expr: ExprRef, ) -> Result, DatabaseError> { self.binder() .bind_unary_op_expr(expr, op, self.arena()) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn function( self, name: impl Into, - args: Vec, + args: Vec, ) -> Result, DatabaseError> { self.binder() .bind_function_call(name.into(), args, self.arena()) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn aggregate( self, kind: AggKind, - args: Vec, + args: Vec, ) -> Result, DatabaseError> { self.binder() .bind_aggregate_function(kind, args, false, self.arena()) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn window( self, kind: WindowFunctionKind, - args: Vec, + args: Vec, spec: WindowSpec, ) -> Result, DatabaseError> { self.binder() .bind_window_function(kind, args, spec.partition_by, spec.order_by, self.arena()) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn scalar_subquery( @@ -555,7 +603,7 @@ where let mut context = OrmContext { binder, arena }; build(&mut context) }) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn exists_subquery( @@ -574,14 +622,14 @@ where let mut context = OrmContext { binder, arena }; build(&mut context) }) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } fn quantified_subquery( self, quantifier: MarkApplyQuantifier, negated: bool, - left_expr: ScalarExpression, + left_expr: ExprRef, compare_op: expression::BinaryOperator, build: F, ) -> Result, DatabaseError> @@ -603,7 +651,7 @@ where build(&mut context) }, ) - .map(|expr| self.wrap(expr)) + .map(|expr| self.alloc(expr)) } } @@ -611,8 +659,9 @@ where /// /// `CtxExpression` is a scope-bound ORM expression handle, not a reusable core /// expression value. It exists so ORM code can use natural chained binding such -/// as `e.column(User::age())?.gte(18)?`. Convert it to a core -/// [`ScalarExpression`] only at ORM binder boundaries with [`Self::into_scalar`]. +/// as `e.column(User::age())?.gte(18)?`. It retains the arena-backed [`ExprRef`] +/// when passed through ORM expression APIs, avoiding cloning and reallocating an +/// already-bound [`ScalarExpression`]. /// /// This type intentionally cannot be sent or shared across threads, and its /// internal scope handle is private. @@ -621,7 +670,7 @@ where T: Transaction, A: AsRef<[(&'static str, DataValue)]>, { - expr: ScalarExpression, + expr: ExprRef, scope: ExprBindScopeHandle<'bind, 'parent, 'arena, T, A>, } @@ -630,7 +679,7 @@ where T: Transaction, A: AsRef<[(&'static str, DataValue)]>, { - pub fn into_scalar(self) -> ScalarExpression { + pub fn into_scalar(self) -> ExprRef { self.expr } @@ -654,83 +703,83 @@ where self.into_sort().nulls_last() } - pub fn eq(self, right: R) -> Result { + pub fn eq(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::Eq, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn ne(self, right: R) -> Result { + pub fn ne(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::NotEq, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn gt(self, right: R) -> Result { + pub fn gt(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::Gt, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn gte(self, right: R) -> Result { + pub fn gte(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::GtEq, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn lt(self, right: R) -> Result { + pub fn lt(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::Lt, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn lte(self, right: R) -> Result { + pub fn lte(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::LtEq, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn like(self, right: R) -> Result { + pub fn like(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::Like(None), - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn not_like(self, right: R) -> Result { + pub fn not_like(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::NotLike(None), - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn and(self, right: R) -> Result { + pub fn and(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::And, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } - pub fn or(self, right: R) -> Result { + pub fn or(self, right: R) -> Result { self.scope.binary( self.expr, expression::BinaryOperator::Or, - right.into_orm_scalar(), + self.scope.bind(right.into_orm_expression()), ) } @@ -743,82 +792,82 @@ where let scope = self.scope; let expr = ScalarExpression::IsNull { negated: false, - expr: Box::new(self.expr), + expr: self.expr, }; - scope.wrap(expr) + scope.alloc(expr) } pub fn is_not_null(self) -> Self { let scope = self.scope; let expr = ScalarExpression::IsNull { negated: true, - expr: Box::new(self.expr), + expr: self.expr, }; - scope.wrap(expr) + scope.alloc(expr) } pub fn in_list(self, values: I) -> Result where I: IntoIterator, - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let scope = self.scope; let expr = ScalarExpression::In { negated: false, - expr: Box::new(self.expr), + expr: self.expr, args: values .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) + .map(|expr| scope.bind(expr.into_orm_expression())) .collect(), }; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn not_in_list(self, values: I) -> Result where I: IntoIterator, - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let scope = self.scope; let expr = ScalarExpression::In { negated: true, - expr: Box::new(self.expr), + expr: self.expr, args: values .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) + .map(|expr| scope.bind(expr.into_orm_expression())) .collect(), }; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn between(self, low: L, high: H) -> Result where - L: IntoOrmScalarExpression, - H: IntoOrmScalarExpression, + L: IntoOrmExpression, + H: IntoOrmExpression, { let scope = self.scope; let expr = ScalarExpression::Between { negated: false, - expr: Box::new(self.expr), - left_expr: Box::new(low.into_orm_scalar()), - right_expr: Box::new(high.into_orm_scalar()), + expr: self.expr, + left_expr: scope.bind(low.into_orm_expression()), + right_expr: scope.bind(high.into_orm_expression()), }; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn not_between(self, low: L, high: H) -> Result where - L: IntoOrmScalarExpression, - H: IntoOrmScalarExpression, + L: IntoOrmExpression, + H: IntoOrmExpression, { let scope = self.scope; let expr = ScalarExpression::Between { negated: true, - expr: Box::new(self.expr), - left_expr: Box::new(low.into_orm_scalar()), - right_expr: Box::new(high.into_orm_scalar()), + expr: self.expr, + left_expr: scope.bind(low.into_orm_expression()), + right_expr: scope.bind(high.into_orm_expression()), }; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn alias(self, alias: impl Into) -> Self { @@ -827,18 +876,17 @@ where scope .binder() .context - .add_alias(None, alias.clone(), self.expr.clone()); + .add_alias(None, alias.clone(), self.expr); let expr = ScalarExpression::Alias { - expr: Box::new(self.expr), + expr: self.expr, alias: AliasType::Name(alias), }; - scope.wrap(expr) + scope.alloc(expr) } pub fn cast(self, ty: LogicalType) -> Result { let scope = self.scope; - ScalarExpression::type_cast(self.expr, Cow::Owned(ty), scope.arena()) - .map(|expr| scope.wrap(expr)) + Ok(scope.wrap(self.expr.type_cast(Cow::Owned(ty), scope.arena())?)) } pub fn function( @@ -847,14 +895,15 @@ where args: impl IntoIterator, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { + let scope = self.scope; let mut args = args .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) + .map(|expr| scope.bind(expr.into_orm_expression())) .collect::>(); args.insert(0, self.expr); - self.scope.function(name, args) + scope.function(name, args) } fn quantified_subquery( @@ -907,7 +956,7 @@ where { fn clone(&self) -> Self { Self { - expr: self.expr.clone(), + expr: self.expr, scope: self.scope, } } @@ -940,14 +989,13 @@ where } } -impl<'bind, 'parent, 'arena, T, A> IntoOrmScalarExpression - for CtxExpression<'bind, 'parent, 'arena, T, A> +impl<'bind, 'parent, 'arena, T, A> IntoOrmExpression for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, A: AsRef<[(&'static str, DataValue)]>, { - fn into_orm_scalar(self) -> ScalarExpression { - self.into_scalar() + fn into_orm_expression(self) -> OrmExpression { + OrmExpression::Bound(self.expr) } } @@ -1025,7 +1073,7 @@ where binder: &'ctx mut Binder<'bind, 'parent, T, A>, arena: &'ctx mut PlanArena<'arena>, source_name: String, - value_exprs: Vec<(ColumnRef, ScalarExpression)>, + value_exprs: Vec<(ColumnRef, ExprRef)>, } impl<'ctx, 'bind, 'parent, 'arena, T, A> OrmContext<'ctx, 'bind, 'parent, 'arena, T, A> @@ -1068,7 +1116,7 @@ where if mutation_source { self.binder.with_pk(source.table_name.as_str().into()); } - let plan = bind_orm_source(self.binder, source.clone(), None, self.arena); + let plan = bind_orm_source(self.binder, source, None, self.arena); if mutation_source { self.binder.clear_with_pk(); } @@ -1258,7 +1306,7 @@ where ExprBindScopeHandle::new(self) } - fn wrap(&self, expr: ScalarExpression) -> CtxExpression<'bind, 'parent, 'arena, T, A> { + fn wrap(&self, expr: impl Into) -> CtxExpression<'bind, 'parent, 'arena, T, A> { self.handle().wrap(expr) } @@ -1273,7 +1321,7 @@ where None, scope.arena(), )?; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn qualified_column( @@ -1288,7 +1336,7 @@ where None, scope.arena(), )?; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } #[doc(hidden)] @@ -1302,7 +1350,7 @@ where scope .binder() .bind_column_ref_by_name(Some(relation), column, None, scope.arena())?; - Ok(scope.wrap(expr)) + Ok(scope.alloc(expr)) } pub fn value(&self, value: V) -> CtxExpression<'bind, 'parent, 'arena, T, A> { @@ -1315,134 +1363,142 @@ where pub fn alias( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, alias: impl Into, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> { - self.wrap(expr.into_orm_scalar()).alias(alias) + self.wrap(expr.into_orm_expression()).alias(alias) } pub fn cast( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, ty: LogicalType, ) -> Result, DatabaseError> { let scope = self.handle(); - let expr = - ScalarExpression::type_cast(expr.into_orm_scalar(), Cow::Owned(ty), scope.arena())?; - Ok(scope.wrap(expr)) + Ok(scope.wrap( + scope + .bind(expr.into_orm_expression()) + .type_cast(Cow::Owned(ty), scope.arena())?, + )) } pub fn unary( &self, op: expression::UnaryOperator, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, ) -> Result, DatabaseError> { - self.handle().unary(op, expr.into_orm_scalar()) + let scope = self.handle(); + scope.unary(op, scope.bind(expr.into_orm_expression())) } pub fn binary( &self, - left: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, op: expression::BinaryOperator, - right: impl IntoOrmScalarExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { - self.handle() - .binary(left.into_orm_scalar(), op, right.into_orm_scalar()) + let scope = self.handle(); + scope.binary( + scope.bind(left.into_orm_expression()), + op, + scope.bind(right.into_orm_expression()), + ) } pub fn eq( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::Eq, right) } pub fn ne( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::NotEq, right) } pub fn gt( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::Gt, right) } pub fn gte( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::GtEq, right) } pub fn lt( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::Lt, right) } pub fn lte( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::LtEq, right) } pub fn and( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::And, right) } pub fn or( &self, - left: impl IntoOrmScalarExpression, - right: impl IntoOrmScalarExpression, + left: impl IntoOrmExpression, + right: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.binary(left, expression::BinaryOperator::Or, right) } pub fn is_null( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> { - self.wrap(expr.into_orm_scalar()).is_null() + self.wrap(expr.into_orm_expression()).is_null() } pub fn is_not_null( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> { - self.wrap(expr.into_orm_scalar()).is_not_null() + self.wrap(expr.into_orm_expression()).is_not_null() } pub fn in_list( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, args: I, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> where I: IntoIterator, - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { + let scope = self.handle(); let expr = ScalarExpression::In { negated: false, - expr: Box::new(expr.into_orm_scalar()), + expr: scope.bind(expr.into_orm_expression()), args: args .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) + .map(|expr| scope.bind(expr.into_orm_expression())) .collect(), }; self.wrap(expr) @@ -1450,19 +1506,20 @@ where pub fn not_in_list( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, args: I, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> where I: IntoIterator, - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { + let scope = self.handle(); let expr = ScalarExpression::In { negated: true, - expr: Box::new(expr.into_orm_scalar()), + expr: scope.bind(expr.into_orm_expression()), args: args .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) + .map(|expr| scope.bind(expr.into_orm_expression())) .collect(), }; self.wrap(expr) @@ -1470,37 +1527,39 @@ where pub fn between( &self, - expr: impl IntoOrmScalarExpression, - low: impl IntoOrmScalarExpression, - high: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, + low: impl IntoOrmExpression, + high: impl IntoOrmExpression, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> { + let scope = self.handle(); let expr = ScalarExpression::Between { negated: false, - expr: Box::new(expr.into_orm_scalar()), - left_expr: Box::new(low.into_orm_scalar()), - right_expr: Box::new(high.into_orm_scalar()), + expr: scope.bind(expr.into_orm_expression()), + left_expr: scope.bind(low.into_orm_expression()), + right_expr: scope.bind(high.into_orm_expression()), }; self.wrap(expr) } pub fn not_between( &self, - expr: impl IntoOrmScalarExpression, - low: impl IntoOrmScalarExpression, - high: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, + low: impl IntoOrmExpression, + high: impl IntoOrmExpression, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> { + let scope = self.handle(); let expr = ScalarExpression::Between { negated: true, - expr: Box::new(expr.into_orm_scalar()), - left_expr: Box::new(low.into_orm_scalar()), - right_expr: Box::new(high.into_orm_scalar()), + expr: scope.bind(expr.into_orm_expression()), + left_expr: scope.bind(low.into_orm_expression()), + right_expr: scope.bind(high.into_orm_expression()), }; self.wrap(expr) } pub fn not( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, ) -> Result, DatabaseError> { self.unary(expression::UnaryOperator::Not, expr) } @@ -1511,13 +1570,15 @@ where args: impl IntoIterator, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { - let args = args - .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) - .collect(); - self.handle().function(name, args) + let scope = self.handle(); + scope.function( + name, + args.into_iter() + .map(|expr| scope.bind(expr.into_orm_expression())) + .collect(), + ) } pub fn aggregate( @@ -1526,24 +1587,27 @@ where args: impl IntoIterator, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { - let args = args - .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) - .collect(); - self.handle().aggregate(kind, args) + let scope = self.handle(); + scope.aggregate( + kind, + args.into_iter() + .map(|expr| scope.bind(expr.into_orm_expression())) + .collect(), + ) } fn aggregate_window( &self, kind: AggKind, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { - self.handle().window( + let scope = self.handle(); + scope.window( WindowFunctionKind::Aggregate(kind), - vec![expr.into_orm_scalar()], + vec![scope.bind(expr.into_orm_expression())], spec, ) } @@ -1574,7 +1638,7 @@ where pub fn count_over( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { self.aggregate_window(AggKind::Count, expr, spec) @@ -1582,7 +1646,7 @@ where pub fn sum_over( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { self.aggregate_window(AggKind::Sum, expr, spec) @@ -1590,7 +1654,7 @@ where pub fn avg_over( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { self.aggregate_window(AggKind::Avg, expr, spec) @@ -1598,7 +1662,7 @@ where pub fn min_over( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { self.aggregate_window(AggKind::Min, expr, spec) @@ -1606,7 +1670,7 @@ where pub fn max_over( &self, - expr: impl IntoOrmScalarExpression, + expr: impl IntoOrmExpression, spec: WindowSpec, ) -> Result, DatabaseError> { self.aggregate_window(AggKind::Max, expr, spec) @@ -1616,17 +1680,19 @@ where &self, spec: WindowSpec, ) -> Result, DatabaseError> { - self.handle().window( + let scope = self.handle(); + scope.window( WindowFunctionKind::Aggregate(AggKind::Count), - vec![Binder::<'bind, 'parent, T, A>::wildcard_expr()], + vec![scope.bind(Binder::<'bind, 'parent, T, A>::wildcard_expr())], spec, ) } pub fn count_all(&self) -> Result, DatabaseError> { - self.aggregate( + let scope = self.handle(); + scope.aggregate( AggKind::Count, - vec![Binder::<'bind, 'parent, T, A>::wildcard_expr()], + vec![scope.bind(Binder::<'bind, 'parent, T, A>::wildcard_expr())], ) } @@ -1636,15 +1702,21 @@ where else_expr: Option, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> where - C: IntoOrmScalarExpression, - V: IntoOrmScalarExpression, - E: IntoOrmScalarExpression, + C: IntoOrmExpression, + V: IntoOrmExpression, + E: IntoOrmExpression, { + let scope = self.handle(); let expr_pairs = expr_pairs .into_iter() - .map(|(condition, value)| (condition.into_orm_scalar(), value.into_orm_scalar())) + .map(|(condition, value)| { + ( + scope.bind(condition.into_orm_expression()), + scope.bind(value.into_orm_expression()), + ) + }) .collect::>(); - let else_expr = else_expr.map(IntoOrmScalarExpression::into_orm_scalar); + let else_expr = else_expr.map(|expr| scope.bind(expr.into_orm_expression())); let ty = expr_pairs .first() .map(|(_, value)| value.return_type(self.arena).into_owned()) @@ -1657,27 +1729,34 @@ where self.wrap(ScalarExpression::CaseWhen { operand_expr: None, expr_pairs, - else_expr: else_expr.map(Box::new), + else_expr, ty, }) } pub fn case_value( &self, - operand_expr: impl IntoOrmScalarExpression, + operand_expr: impl IntoOrmExpression, expr_pairs: impl IntoIterator, else_expr: Option, ) -> CtxExpression<'bind, 'parent, 'arena, T, A> where - K: IntoOrmScalarExpression, - V: IntoOrmScalarExpression, - E: IntoOrmScalarExpression, + K: IntoOrmExpression, + V: IntoOrmExpression, + E: IntoOrmExpression, { + let scope = self.handle(); + let operand_expr = scope.bind(operand_expr.into_orm_expression()); let expr_pairs = expr_pairs .into_iter() - .map(|(key, value)| (key.into_orm_scalar(), value.into_orm_scalar())) + .map(|(key, value)| { + ( + scope.bind(key.into_orm_expression()), + scope.bind(value.into_orm_expression()), + ) + }) .collect::>(); - let else_expr = else_expr.map(IntoOrmScalarExpression::into_orm_scalar); + let else_expr = else_expr.map(|expr| scope.bind(expr.into_orm_expression())); let ty = expr_pairs .first() .map(|(_, value)| value.return_type(self.arena).into_owned()) @@ -1688,9 +1767,9 @@ where }) .unwrap_or(LogicalType::SqlNull); self.wrap(ScalarExpression::CaseWhen { - operand_expr: Some(Box::new(operand_expr.into_orm_scalar())), + operand_expr: Some(operand_expr), expr_pairs, - else_expr: else_expr.map(Box::new), + else_expr, ty, }) } @@ -1732,7 +1811,9 @@ where where D: ToDataValue, { - let expr = ScalarExpression::Constant(value.to_data_value()); + let expr = self + .arena + .alloc_expression(ScalarExpression::Constant(value.to_data_value())); self.push_assignment(field.column, expr) } @@ -1766,14 +1847,15 @@ where ) -> Result, ) -> Result<(), DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let expr = with_query_bind_step!(self.binder, QueryBindStep::Project, { let mut scope = ExprBindScope { binder: self.binder, arena: self.arena, }; - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); self.push_assignment(field.column, expr?) } @@ -1781,21 +1863,21 @@ where fn push_assignment( &mut self, column_name: &str, - mut expr: ScalarExpression, + mut expr: ExprRef, ) -> Result<(), DatabaseError> { let column = bind_orm_target_column(self.binder, &self.source_name, column_name, self.arena)?; - if matches!(expr, ScalarExpression::Empty) { + if matches!(self.arena.expression(expr), ScalarExpression::Empty) { let column_catalog = self.arena.column(column); let default_value = column_catalog - .default_value()? + .default_value(self.arena)? .ok_or(DatabaseError::DefaultNotExist)?; - expr = ScalarExpression::Constant(default_value); + expr = self + .arena + .alloc_expression(ScalarExpression::Constant(default_value)); } - let column_catalog = self.arena.column(column); - expr = ScalarExpression::type_cast( - expr, - Cow::Borrowed(column_catalog.datatype()), + expr = expr.type_cast( + Cow::Owned(self.arena.column(column).datatype().clone()), self.arena, )?; self.value_exprs.push((column, expr)); @@ -1869,11 +1951,12 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let predicate = with_query_bind_step!(self.binder, QueryBindStep::Where, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); let predicate = predicate?; self.filter_expr(predicate) @@ -1911,7 +1994,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let source = match alias { Some(alias) => QuerySource::model::().with_alias(alias), @@ -1930,7 +2013,8 @@ where self.binder.extend(right_context); let on = with_query_bind_step!(self.binder, QueryBindStep::From, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); self.plan = self.binder.bind_join_plans( self.plan, @@ -1949,7 +2033,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::Inner, None, build) } @@ -1962,7 +2046,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::Inner, Some(alias.into()), build) } @@ -1974,7 +2058,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::LeftOuter, None, build) } @@ -1987,7 +2071,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::LeftOuter, Some(alias.into()), build) } @@ -1999,7 +2083,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::RightOuter, None, build) } @@ -2011,7 +2095,7 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.join_on::(JoinType::Full, None, build) } @@ -2075,7 +2159,7 @@ where &relation, Field::::new(M::table_name(), field.column), )? - .into_orm_scalar(), + .into_scalar(), ); } })?; @@ -2101,11 +2185,12 @@ where ) -> Result, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let expr = with_query_bind_step!(self.binder, QueryBindStep::Project, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); Ok(self.select_list(vec![expr?])) } @@ -2117,18 +2202,17 @@ where ) -> Result, DatabaseError>, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let exprs = with_query_bind_step!(self.binder, QueryBindStep::Project, { let mut scope = self.expr_scope(); + let handle = scope.handle(); build(&mut scope)? - }); - Ok(self.select_list( - exprs? .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) - .collect(), - )) + .map(|expr| handle.bind(expr.into_orm_expression())) + .collect::>() + }); + Ok(self.select_list(exprs?)) } pub fn project_scalar( @@ -2162,7 +2246,7 @@ where ) -> Result, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.project_model()?.group_by(build) } @@ -2174,7 +2258,7 @@ where ) -> Result, ) -> Result, DatabaseError> where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { self.project_model()?.having(build) } @@ -2217,7 +2301,7 @@ where let count = with_query_bind_step!(self.binder, QueryBindStep::Project, { let scope = self.expr_scope(); let count = scope.count_all()?; - scope.alias(count, "count").into_orm_scalar() + scope.alias(count, "count").into_scalar() }); self.select_list(vec![count?]).count() } @@ -2287,11 +2371,12 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let expr = with_query_bind_step!(self.binder, QueryBindStep::Project, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); Ok(self.set_select_list(vec![expr?])) } @@ -2303,18 +2388,17 @@ where ) -> Result, DatabaseError>, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let exprs = with_query_bind_step!(self.binder, QueryBindStep::Project, { let mut scope = self.expr_scope(); + let handle = scope.handle(); build(&mut scope)? - }); - Ok(self.set_select_list( - exprs? .into_iter() - .map(IntoOrmScalarExpression::into_orm_scalar) - .collect(), - )) + .map(|expr| handle.bind(expr.into_orm_expression())) + .collect::>() + }); + Ok(self.set_select_list(exprs?)) } pub fn project_scalar( @@ -2346,11 +2430,12 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let expr = with_query_bind_step!(self.binder, QueryBindStep::Agg, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); self.group_by_expr(expr?) } @@ -2362,11 +2447,12 @@ where ) -> Result, ) -> Result where - E: IntoOrmScalarExpression, + E: IntoOrmExpression, { let expr = with_query_bind_step!(self.binder, QueryBindStep::Having, { let mut scope = self.expr_scope(); - build(&mut scope)?.into_orm_scalar() + let handle = scope.handle(); + handle.bind(build(&mut scope)?.into_orm_expression()) }); self.having_expr(expr?) } @@ -2421,7 +2507,7 @@ where let count = with_query_bind_step!(self.binder, QueryBindStep::Project, { let scope = self.expr_scope(); let count = scope.count_all()?; - scope.alias(count, "count").into_orm_scalar() + scope.alias(count, "count").into_scalar() }); self.set_select_list(vec![count?]) .aggregate_without_group()? @@ -2434,7 +2520,7 @@ pub trait Projection: FromQueryRow { fn bind_projection<'ctx, 'bind, 'parent, 'arena, T, A>( scope: &mut ExprBindScope<'ctx, 'bind, 'parent, 'arena, T, A>, relation: &str, - ) -> Result, DatabaseError> + ) -> Result, DatabaseError> where T: Transaction, A: AsRef<[(&'static str, DataValue)]>; @@ -2511,12 +2597,15 @@ where .iter() .copied() .enumerate() - .map(|(position, target_column)| ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( + .map(|(position, target_column)| { + let expr = arena.alloc_expression(ScalarExpression::column_expr( input_schema[position], position, - )), - alias: AliasType::Name(arena.column(target_column).name().to_string()), + )); + arena.alloc_expression(ScalarExpression::Alias { + expr, + alias: AliasType::Name(arena.column(target_column).name().to_string()), + }) }) .collect::>() } else { @@ -2528,13 +2617,14 @@ where let column = source .column(&column_name, arena) .ok_or_else(|| DatabaseError::column_not_found(column_name.clone()))?; - projection.push(ScalarExpression::Alias { - expr: Box::new(ScalarExpression::column_expr( - input_schema[position], - position, - )), + let expr = arena.alloc_expression(ScalarExpression::column_expr( + input_schema[position], + position, + )); + projection.push(arena.alloc_expression(ScalarExpression::Alias { + expr, alias: AliasType::Name(arena.column(column).name().to_string()), - }); + })); } projection } @@ -2618,8 +2708,8 @@ pub trait Model: Sized + FromQueryRow { /// /// `#[derive(Model)]` generates this automatically. Manual implementations /// can override it to opt into [`Database::migrate`](crate::orm::Database::migrate). - fn columns() -> &'static [ColumnCatalog] { - &[] + fn columns(_arena: &mut crate::planner::TableArena) -> Vec { + Vec::new() } /// Returns secondary indexes declared by the model. @@ -3120,12 +3210,18 @@ impl_from_query_tuple!( (A, B, C, D, E, F, G, H), ); -fn model_column_default(model: &ColumnCatalog) -> Result, DatabaseError> { - model.default_value() +fn model_column_default( + model: &ColumnCatalog, + arena: &PlanArena<'_>, +) -> Result, DatabaseError> { + model.default_value(arena) } -fn catalog_column_default(column: &ColumnCatalog) -> Result, DatabaseError> { - column.default_value() +fn catalog_column_default( + column: &ColumnCatalog, + arena: &PlanArena<'_>, +) -> Result, DatabaseError> { + column.default_value(arena) } fn model_column_type_matches_catalog(model: &ColumnCatalog, column: &ColumnCatalog) -> bool { @@ -3135,23 +3231,25 @@ fn model_column_type_matches_catalog(model: &ColumnCatalog, column: &ColumnCatal fn model_column_matches_catalog( model: &ColumnCatalog, column: &ColumnCatalog, + arena: &PlanArena<'_>, ) -> Result { Ok(model.desc().is_primary() == column.desc().is_primary() && model.desc().is_unique() == column.desc().is_unique() && model.nullable() == column.nullable() && model_column_type_matches_catalog(model, column) - && model_column_default(model)? == catalog_column_default(column)?) + && model_column_default(model, arena)? == catalog_column_default(column, arena)?) } fn model_column_rename_compatible( model: &ColumnCatalog, column: &ColumnCatalog, + arena: &PlanArena<'_>, ) -> Result { Ok(model.desc().is_primary() == column.desc().is_primary() && model.desc().is_unique() == column.desc().is_unique() && model.nullable() == column.nullable() && model_column_type_matches_catalog(model, column) - && model_column_default(model)? == catalog_column_default(column)?) + && model_column_default(model, arena)? == catalog_column_default(column, arena)?) } fn extract_optional_model(iter: I) -> Result, DatabaseError> @@ -3763,9 +3861,9 @@ mod tests { assert_eq!( plan, concat!( - "Projection [#2] [Project => (Sort Option: Follow)] ", - "Filter (#1 >= 4), Is Having: false [Filter => (Sort Option: Follow)] ", - "TableScan orm_unit_users -> [#1, #2] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.name] [Project => (Sort Option: Follow)] ", + "Filter (orm_unit_users.id >= 4), Is Having: false [Filter => (Sort Option: Follow)] ", + "TableScan orm_unit_users -> [orm_unit_users.id, orm_unit_users.name] [SeqScan => (Sort Option: None)]" ), "{plan}" ); @@ -3801,9 +3899,9 @@ mod tests { expression_plan, concat!( "Projection [upper_name] [Project => (Sort Option: Follow)] ", - "Sort By #3 Desc Nulls Last [Sort => (Sort Option: OrderBy: (#3 Desc Nulls Last) ignore_prefix_len: 0)] ", - "Filter ((#3 is not null && ((#3 >= 18) && (#3 <= 25))) && (!(#2 != Bob) && (#2 = Missing))), Is Having: false ", - "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [#2, #3] [SeqScan => (Sort Option: None)]" + "Sort By orm_unit_users.age Desc Nulls Last [Sort => (Sort Option: OrderBy: (orm_unit_users.age Desc Nulls Last) ignore_prefix_len: 0)] ", + "Filter ((orm_unit_users.age is not null && ((orm_unit_users.age >= 18) && (orm_unit_users.age <= 25))) && (!(orm_unit_users.name != Bob) && (orm_unit_users.name = Missing))), Is Having: false ", + "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [orm_unit_users.name, orm_unit_users.age] [SeqScan => (Sort Option: None)]" ), "{expression_plan}" ); @@ -3817,16 +3915,16 @@ mod tests { in_list.and(not_in_list) })? .project_scalar(OrmUnitUser::id())? - .order_by(SortField::from(ScalarExpression::from(1_i32)).asc())? + .order_by_expr(|e| Ok(e.value(1_i32).asc()))? .finish() })?; assert_eq!( list_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", "Sort By 1 Asc Nulls Last [Sort => (Sort Option: OrderBy: (1 Asc Nulls Last) ignore_prefix_len: 0)] ", - "Filter (((#1 = 3) || ((#1 = 2) || (#1 = 1))) && (#1 != 3)), Is Having: false ", - "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)]" + "Filter (((orm_unit_users.id = 3) || ((orm_unit_users.id = 2) || (orm_unit_users.id = 1))) && (orm_unit_users.id != 3)), Is Having: false ", + "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)]" ), "{list_plan}" ); @@ -3845,9 +3943,9 @@ mod tests { assert_eq!( nullable_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "Filter (((#3 < 10) || (#3 > 30)) || #3 is null), Is Having: false ", - "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [#1, #3] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "Filter (((orm_unit_users.age < 10) || (orm_unit_users.age > 30)) || orm_unit_users.age is null), Is Having: false ", + "[Filter => (Sort Option: Follow)] TableScan orm_unit_users -> [orm_unit_users.id, orm_unit_users.age] [SeqScan => (Sort Option: None)]" ), "{nullable_plan}" ); @@ -3878,11 +3976,11 @@ mod tests { assert_eq!( grouped_plan, concat!( - "Projection [#5, #7] [Project => (Sort Option: Follow)] ", - "Sort By #5 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#5 Asc Nulls Last) ignore_prefix_len: 0)] ", - "Filter (Sum(#6) >= 200), Is Having: true [Filter => (Sort Option: Follow)] ", - "Aggregate [Sum(#6)] -> Group By [#5] [HashAggregate => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#5, #6] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_orders.user_id, Sum(orm_unit_orders.amount)] [Project => (Sort Option: Follow)] ", + "Sort By orm_unit_orders.user_id Asc Nulls Last [Sort => (Sort Option: OrderBy: (orm_unit_orders.user_id Asc Nulls Last) ignore_prefix_len: 0)] ", + "Filter (Sum(orm_unit_orders.amount) >= 200), Is Having: true [Filter => (Sort Option: Follow)] ", + "Aggregate [Sum(orm_unit_orders.amount)] -> Group By [orm_unit_orders.user_id] [HashAggregate => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id, orm_unit_orders.amount] [SeqScan => (Sort Option: None)]" ), "{grouped_plan}" ); @@ -3899,10 +3997,10 @@ mod tests { assert_eq!( right_join_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "RightOuter Join On #1 = #5 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#5] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "RightOuter Join On orm_unit_users.id = orm_unit_orders.user_id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id] [SeqScan => (Sort Option: None)]" ), "{right_join_plan}" ); @@ -3919,10 +4017,10 @@ mod tests { assert_eq!( full_join_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "Full Join On #1 = #5 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#5] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "Full Join On orm_unit_users.id = orm_unit_orders.user_id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id] [SeqScan => (Sort Option: None)]" ), "{full_join_plan}" ); @@ -3936,10 +4034,10 @@ mod tests { assert_eq!( cross_join_plan, concat!( - "Projection [#1, #4] [Project => (Sort Option: Follow)] ", + "Projection [orm_unit_users.id, orm_unit_orders.id] [Project => (Sort Option: Follow)] ", "Cross Join Nothing [NestLoopJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#4] [SeqScan => (Sort Option: None)]" + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.id] [SeqScan => (Sort Option: None)]" ), "{cross_join_plan}" ); @@ -3953,10 +4051,10 @@ mod tests { assert_eq!( inner_using_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "Inner Join On #1 = #4 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#4] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "Inner Join On orm_unit_users.id = orm_unit_orders.id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.id] [SeqScan => (Sort Option: None)]" ), "{inner_using_plan}" ); @@ -3970,10 +4068,10 @@ mod tests { assert_eq!( left_using_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "LeftOuter Join On #1 = #4 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#4] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "LeftOuter Join On orm_unit_users.id = orm_unit_orders.id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.id] [SeqScan => (Sort Option: None)]" ), "{left_using_plan}" ); @@ -3987,10 +4085,10 @@ mod tests { assert_eq!( right_using_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "RightOuter Join On #1 = #4 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#4] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "RightOuter Join On orm_unit_users.id = orm_unit_orders.id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.id] [SeqScan => (Sort Option: None)]" ), "{right_using_plan}" ); @@ -4004,10 +4102,10 @@ mod tests { assert_eq!( full_using_plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "Full Join On #1 = #4 [HashJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#4] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "Full Join On orm_unit_users.id = orm_unit_orders.id [HashJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.id] [SeqScan => (Sort Option: None)]" ), "{full_using_plan}" ); @@ -4032,10 +4130,10 @@ mod tests { assert_eq!( plan, concat!( - "Projection [#1] [Project => (Sort Option: Follow)] ", - "Inner Join On #1 = #5 [NestLoopJoin => (Sort Option: None)] ", - "TableScan orm_unit_users -> [#1] [SeqScan => (Sort Option: None)] ", - "TableScan orm_unit_orders -> [#5] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_users.id] [Project => (Sort Option: Follow)] ", + "Inner Join On orm_unit_users.id = orm_unit_orders.user_id [NestLoopJoin => (Sort Option: None)] ", + "TableScan orm_unit_users -> [orm_unit_users.id] [SeqScan => (Sort Option: None)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id] [SeqScan => (Sort Option: None)]" ), "{plan}" ); @@ -4064,10 +4162,10 @@ mod tests { assert_eq!( plan, concat!( - "Projection [#5, #7] [Project => (Sort Option: Follow)] ", - "Aggregate [Sum(#6)] -> Group By [#5] [StreamAggregate => (Sort Option: Follow)] ", - "Sort By #5 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#5 Asc Nulls Last) ignore_prefix_len: 0)] ", - "TableScan orm_unit_orders -> [#5, #6] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_orders.user_id, Sum(orm_unit_orders.amount)] [Project => (Sort Option: Follow)] ", + "Aggregate [Sum(orm_unit_orders.amount)] -> Group By [orm_unit_orders.user_id] [StreamAggregate => (Sort Option: Follow)] ", + "Sort By orm_unit_orders.user_id Asc Nulls Last [Sort => (Sort Option: OrderBy: (orm_unit_orders.user_id Asc Nulls Last) ignore_prefix_len: 0)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id, orm_unit_orders.amount] [SeqScan => (Sort Option: None)]" ), "{plan}" ); @@ -4082,10 +4180,10 @@ mod tests { assert_eq!( distinct_plan, concat!( - "Projection [#5] [Project => (Sort Option: Follow)] ", - "Aggregate [] -> Group By [#5] [StreamDistinct => (Sort Option: Follow)] ", - "Sort By #5 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#5 Asc Nulls Last) ignore_prefix_len: 0)] ", - "TableScan orm_unit_orders -> [#5] [SeqScan => (Sort Option: None)]" + "Projection [orm_unit_orders.user_id] [Project => (Sort Option: Follow)] ", + "Aggregate [] -> Group By [orm_unit_orders.user_id] [StreamDistinct => (Sort Option: Follow)] ", + "Sort By orm_unit_orders.user_id Asc Nulls Last [Sort => (Sort Option: OrderBy: (orm_unit_orders.user_id Asc Nulls Last) ignore_prefix_len: 0)] ", + "TableScan orm_unit_orders -> [orm_unit_orders.user_id] [SeqScan => (Sort Option: None)]" ), "{distinct_plan}" ); diff --git a/src/planner/arena.rs b/src/planner/arena.rs index c041281a..c17a709f 100644 --- a/src/planner/arena.rs +++ b/src/planner/arena.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::catalog::{ColumnCatalog, ColumnRef, TableName}; +use crate::expression::ScalarExpression; use crate::planner::LogicalPlan; use crate::types::index::{IndexMeta, IndexMetaRef}; use crate::types::tuple::Schema; @@ -24,6 +25,7 @@ pub struct TableArena { dummy_columns: [ColumnCatalog; DUMMY_COLUMN_COUNT], columns: Vec, indexes: Vec, + expressions: Vec, version: usize, } @@ -37,6 +39,11 @@ struct TableArenaIndex { live: bool, } +struct TableArenaExpression { + expression: ScalarExpression, + live: bool, +} + pub struct TableArenaCell { value: UnsafeCell, } @@ -55,19 +62,43 @@ pub struct PlanArena<'a> { temp_table_id: usize, columns: Vec, indexes: Vec, + expressions: Vec, plans: Vec, } +#[derive(Debug, Clone, Copy, Hash, Eq, PartialEq)] +pub struct ExprRef { + pos: usize, +} + #[derive(Debug, Clone, Copy, Hash, Eq, PartialEq)] pub(crate) struct PlanRef { pos: usize, } +impl ExprRef { + pub(crate) fn new(pos: usize) -> Self { + Self { pos } + } + + pub(crate) fn pos(self) -> usize { + self.pos + } +} + +impl fmt::Display for ExprRef { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "#expr{}", self.pos) + } +} + pub trait MetaArena { fn alloc_column(&mut self, column: ColumnCatalog) -> ColumnRef; fn alloc_index(&mut self, index: IndexMeta) -> IndexMetaRef; + fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef; + fn alloc_columns(&mut self, columns: I) -> Schema where Self: Sized, @@ -83,6 +114,8 @@ pub trait MetaArena { fn index(&self, index: IndexMetaRef) -> &IndexMeta; + fn expression(&self, expression: ExprRef) -> &ScalarExpression; + fn find_column(&self, column: &ColumnCatalog) -> Option; fn find_index(&self, index: &IndexMeta) -> Option; @@ -150,6 +183,7 @@ impl Default for TableArena { }), columns: Vec::new(), indexes: Vec::new(), + expressions: Vec::new(), version: 0, } } @@ -178,6 +212,10 @@ impl TableArena { ::alloc_index(self, index) } + pub fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + ::alloc_expression(self, expression) + } + pub(crate) fn column(&self, column: ColumnRef) -> &ColumnCatalog { ::column(self, column) } @@ -186,6 +224,10 @@ impl TableArena { ::index(self, index) } + pub(crate) fn expression(&self, expression: ExprRef) -> &ScalarExpression { + ::expression(self, expression) + } + fn dummy_column(&self, column: ColumnRef) -> Option<&ColumnCatalog> { column .pos() @@ -301,6 +343,29 @@ impl MetaArena for TableArena { IndexMetaRef::new(pos) } + fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + if let Some((pos, slot)) = self + .expressions + .iter_mut() + .enumerate() + .find(|(_, expression)| !expression.live) + { + *slot = TableArenaExpression { + expression, + live: true, + }; + self.increment_version(); + return ExprRef::new(pos); + } + let pos = self.expressions.len(); + self.expressions.push(TableArenaExpression { + expression, + live: true, + }); + self.increment_version(); + ExprRef::new(pos) + } + fn column(&self, column: ColumnRef) -> &ColumnCatalog { if let Some(column) = self.dummy_column(column) { return column; @@ -320,6 +385,12 @@ impl MetaArena for TableArena { &index.meta } + fn expression(&self, expression: ExprRef) -> &ScalarExpression { + let expression = &self.expressions[expression.pos()]; + assert!(expression.live, "accessing recycled TableArena expression"); + &expression.expression + } + fn find_column(&self, column: &ColumnCatalog) -> Option { self.columns .iter() @@ -347,6 +418,7 @@ impl<'a> PlanArena<'a> { temp_table_id: 0, columns: Vec::new(), indexes: Vec::new(), + expressions: Vec::new(), plans: Vec::new(), } } @@ -379,7 +451,13 @@ impl<'a> PlanArena<'a> { live: true, }); } - if !self.columns.is_empty() || !self.indexes.is_empty() { + for expression in &self.expressions { + table_arena.expressions.push(TableArenaExpression { + expression: expression.clone(), + live: true, + }); + } + if !self.columns.is_empty() || !self.indexes.is_empty() || !self.expressions.is_empty() { table_arena.increment_version(); } } @@ -454,6 +532,24 @@ impl<'a> PlanArena<'a> { &self.plans[plan_ref.pos] } + pub(crate) fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + ::alloc_expression(self, expression) + } + + pub(crate) fn expression(&self, expression_ref: ExprRef) -> &ScalarExpression { + ::expression(self, expression_ref) + } + + pub(crate) fn expression_mut(&mut self, expression_ref: ExprRef) -> &mut ScalarExpression { + self.assert_table_arena_unchanged(); + let persistent_len = self.table_arena.borrow().expressions.len(); + assert!( + expression_ref.pos() >= persistent_len, + "persistent expressions are immutable" + ); + &mut self.expressions[expression_ref.pos() - persistent_len] + } + pub fn column(&self, column: ColumnRef) -> &ColumnCatalog { ::column(self, column) } @@ -490,6 +586,13 @@ impl MetaArena for PlanArena<'_> { IndexMetaRef::new(pos) } + fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + self.assert_table_arena_unchanged(); + let pos = self.table_arena.borrow().expressions.len() + self.expressions.len(); + self.expressions.push(expression); + ExprRef::new(pos) + } + fn column(&self, column: ColumnRef) -> &ColumnCatalog { self.assert_table_arena_unchanged(); let table_arena = self.table_arena.borrow(); @@ -514,6 +617,16 @@ impl MetaArena for PlanArena<'_> { &self.indexes[index.pos() - table_indexes_len] } } + fn expression(&self, expression: ExprRef) -> &ScalarExpression { + self.assert_table_arena_unchanged(); + let table_arena = self.table_arena.borrow(); + let persistent_len = table_arena.expressions.len(); + if expression.pos() < persistent_len { + table_arena.expression(expression) + } else { + &self.expressions[expression.pos() - persistent_len] + } + } fn find_column(&self, column: &ColumnCatalog) -> Option { self.assert_table_arena_unchanged(); diff --git a/src/planner/mod.rs b/src/planner/mod.rs index 99f811d4..a80b0571 100644 --- a/src/planner/mod.rs +++ b/src/planner/mod.rs @@ -23,10 +23,48 @@ use crate::planner::operator::union::UnionOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::{Operator, PhysicalOption}; use kite_sql_serde_macros::ReferenceSerialization; +use std::fmt; use std::hash::{Hash, Hasher}; pub(crate) use arena::PlanRef; -pub use arena::{MetaArena, PlanArena, TableArena, TableArenaCell}; +pub use arena::{ExprRef, MetaArena, PlanArena, TableArena, TableArenaCell}; + +pub(crate) trait Explain { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result; + + fn explain<'a, 'p>(&'a self, arena: &'a PlanArena<'p>) -> ExplainDisplay<'a, 'p, Self> + where + Self: Sized, + { + ExplainDisplay { value: self, arena } + } +} + +pub(crate) struct ExplainDisplay<'a, 'p, T: ?Sized> { + value: &'a T, + arena: &'a PlanArena<'p>, +} + +impl fmt::Display for ExplainDisplay<'_, '_, T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.value.fmt(self.arena, f) + } +} + +pub(crate) fn fmt_explain_list( + values: &[T], + separator: &str, + arena: &PlanArena<'_>, + f: &mut fmt::Formatter<'_>, +) -> fmt::Result { + for (index, value) in values.iter().enumerate() { + if index > 0 { + f.write_str(separator)?; + } + value.fmt(arena, f)?; + } + Ok(()) +} #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub enum Childrens { @@ -306,21 +344,29 @@ impl LogicalPlan { } } - #[allow(clippy::only_used_in_recursion)] pub fn explain(&self, arena: &mut PlanArena, indentation: usize) -> String { - let mut result = format!("{:indent$}{}", "", self.operator, indent = indentation); + format!( + "{:indent$}{}", + "", + Explain::explain(self, arena), + indent = indentation + ) + } +} + +impl Explain for LogicalPlan { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.operator.fmt(arena, f)?; if let Some(physical_option) = &self.physical_option { - result.push_str(&format!(" [{physical_option}]")); + write!(f, " [{}]", physical_option.explain(arena))?; } for child in self.childrens.iter() { - let child = child.explain(arena, indentation + 2); - result.push(' '); - result.push_str(child.trim_start()); + write!(f, " {}", Explain::explain(child, arena))?; } - result + Ok(()) } } @@ -475,11 +521,16 @@ mod tests { left: Box::new(left), right: Box::new(right), }; - let operators = twins - .iter() - .map(|plan| format!("{}", plan.operator)) - .collect::>(); - assert_eq!(operators, vec!["Show Tables", "Show Views"]); + let mut children = twins.iter(); + assert!(matches!( + children.next().map(|plan| &plan.operator), + Some(Operator::ShowTable) + )); + assert!(matches!( + children.next().map(|plan| &plan.operator), + Some(Operator::ShowView) + )); + assert!(children.next().is_none()); let (left, right) = twins.pop_twins(); assert!(matches!(left.operator, Operator::ShowTable)); @@ -503,7 +554,7 @@ mod tests { assert_eq!(tables, vec!["users"]); assert_eq!( plan.explain(&mut arena, 0), - "Limit 5, Offset 2 [Limit => (Sort Option: Follow)] TableScan users -> [#0]" + "Limit 5, Offset 2 [Limit => (Sort Option: Follow)] TableScan users -> [id]" ); let _ = plan.output_schema(&mut arena); diff --git a/src/planner/operator/aggregate.rs b/src/planner/operator/aggregate.rs index 899d9c02..09efe78a 100644 --- a/src/planner/operator/aggregate.rs +++ b/src/planner/operator/aggregate.rs @@ -12,17 +12,14 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::iter_ext::Itertools; -use crate::planner::{Childrens, LogicalPlan}; -use crate::{expression::ScalarExpression, planner::operator::Operator}; +use crate::planner::operator::Operator; +use crate::planner::{fmt_explain_list, Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct AggregateOperator { - pub groupby_exprs: Vec, - pub agg_calls: Vec, + pub groupby_exprs: Vec, + pub agg_calls: Vec, pub is_distinct: bool, pub force_spill: bool, } @@ -30,8 +27,8 @@ pub struct AggregateOperator { impl AggregateOperator { pub fn build( children: LogicalPlan, - agg_calls: Vec, - groupby_exprs: Vec, + agg_calls: Vec, + groupby_exprs: Vec, is_distinct: bool, force_spill: bool, ) -> LogicalPlan { @@ -47,24 +44,16 @@ impl AggregateOperator { } } -impl fmt::Display for AggregateOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let calls = self - .agg_calls - .iter() - .map(|call| format!("{call}")) - .join(", "); - write!(f, "Aggregate [{calls}]")?; - +impl Explain for AggregateOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Aggregate [")?; + fmt_explain_list(&self.agg_calls, ", ", arena, f)?; + f.write_str("]")?; if !self.groupby_exprs.is_empty() { - let groupbys = self - .groupby_exprs - .iter() - .map(|groupby| format!("{groupby}")) - .join(", "); - write!(f, " -> Group By [{groupbys}]")?; + f.write_str(" -> Group By [")?; + fmt_explain_list(&self.groupby_exprs, ", ", arena, f)?; + f.write_str("]")?; } - Ok(()) } } diff --git a/src/planner/operator/alter_table/change_column.rs b/src/planner/operator/alter_table/change_column.rs index 18a6ce1f..f3bae4ce 100644 --- a/src/planner/operator/alter_table/change_column.rs +++ b/src/planner/operator/alter_table/change_column.rs @@ -13,16 +13,14 @@ // limitations under the License. use crate::catalog::TableName; -use crate::expression::ScalarExpression; +use crate::planner::{Explain, ExprRef, PlanArena}; use crate::types::LogicalType; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub enum DefaultChange { NoChange, - Set(ScalarExpression), + Set(ExprRef), Drop, } @@ -43,17 +41,18 @@ pub struct ChangeColumnOperator { pub not_null_change: NotNullChange, } -impl fmt::Display for ChangeColumnOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { +impl Explain for ChangeColumnOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!( f, - "Change {} -> {}.{} ({}, {:?}, {:?})", - self.old_column_name, - self.table_name, - self.new_column_name, - self.data_type, - self.default_change, - self.not_null_change - ) + "Change {} -> {}.{} ({}, ", + self.old_column_name, self.table_name, self.new_column_name, self.data_type + )?; + match &self.default_change { + DefaultChange::NoChange => f.write_str("NoChange")?, + DefaultChange::Set(expr) => write!(f, "Set({})", expr.explain(arena))?, + DefaultChange::Drop => f.write_str("Drop")?, + } + write!(f, ", {:?})", self.not_null_change) } } diff --git a/src/planner/operator/analyze.rs b/src/planner/operator/analyze.rs index 7662dda9..314b4d92 100644 --- a/src/planner/operator/analyze.rs +++ b/src/planner/operator/analyze.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::catalog::TableName; +use crate::planner::{fmt_explain_list, Explain, PlanArena}; use crate::types::index::IndexMetaRef; use kite_sql_serde_macros::ReferenceSerialization; @@ -22,3 +23,11 @@ pub struct AnalyzeOperator { pub index_metas: Vec, pub histogram_buckets: Option, } + +impl Explain for AnalyzeOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Analyze {} -> [", self.table_name)?; + fmt_explain_list(&self.index_metas, ", ", arena, f)?; + f.write_str("]") + } +} diff --git a/src/planner/operator/copy_from_file.rs b/src/planner/operator/copy_from_file.rs index 51836e59..2c28a6db 100644 --- a/src/planner/operator/copy_from_file.rs +++ b/src/planner/operator/copy_from_file.rs @@ -14,11 +14,9 @@ use crate::binder::copy::ExtSource; use crate::catalog::TableName; -use crate::iter_ext::Itertools; +use crate::planner::{fmt_explain_list, Explain, PlanArena}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct CopyFromFileOperator { @@ -27,17 +25,10 @@ pub struct CopyFromFileOperator { pub schema_ref: Schema, } -impl fmt::Display for CopyFromFileOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let columns = self.schema_ref.iter().join(", "); - write!( - f, - "Copy {} -> {} [{}]", - self.source.path.display(), - self.table, - columns - )?; - - Ok(()) +impl Explain for CopyFromFileOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Copy {} -> {} [", self.source.path.display(), self.table)?; + fmt_explain_list(&self.schema_ref, ", ", arena, f)?; + f.write_str("]") } } diff --git a/src/planner/operator/create_index.rs b/src/planner/operator/create_index.rs index 92622660..638926e0 100644 --- a/src/planner/operator/create_index.rs +++ b/src/planner/operator/create_index.rs @@ -13,11 +13,9 @@ // limitations under the License. use crate::catalog::{ColumnRef, TableName}; -use crate::iter_ext::Itertools; +use crate::planner::{fmt_explain_list, Explain, PlanArena}; use crate::types::index::IndexType; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct CreateIndexOperator { @@ -29,15 +27,10 @@ pub struct CreateIndexOperator { pub ty: IndexType, } -impl fmt::Display for CreateIndexOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let columns = self.columns.iter().join(", "); - write!( - f, - "Create Index On {} -> [{}], If Not Exists: {}", - self.table_name, columns, self.if_not_exists - )?; - - Ok(()) +impl Explain for CreateIndexOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Create Index On {} -> [", self.table_name)?; + fmt_explain_list(&self.columns, ", ", arena, f)?; + write!(f, "], If Not Exists: {}", self.if_not_exists) } } diff --git a/src/planner/operator/filter.rs b/src/planner/operator/filter.rs index 85135184..8da2499f 100644 --- a/src/planner/operator/filter.rs +++ b/src/planner/operator/filter.rs @@ -12,23 +12,20 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::expression::ScalarExpression; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; use super::Operator; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct FilterOperator { - pub predicate: ScalarExpression, + pub predicate: ExprRef, pub is_optimized: bool, pub having: bool, } impl FilterOperator { - pub fn build(predicate: ScalarExpression, children: LogicalPlan, having: bool) -> LogicalPlan { + pub fn build(predicate: ExprRef, children: LogicalPlan, having: bool) -> LogicalPlan { LogicalPlan::new( Operator::Filter(FilterOperator { predicate, @@ -40,10 +37,13 @@ impl FilterOperator { } } -impl fmt::Display for FilterOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - write!(f, "Filter {}, Is Having: {}", self.predicate, self.having)?; - - Ok(()) +impl Explain for FilterOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "Filter {}, Is Having: {}", + self.predicate.explain(arena), + self.having + ) } } diff --git a/src/planner/operator/join.rs b/src/planner/operator/join.rs index 0a2e671c..8580eae4 100644 --- a/src/planner/operator/join.rs +++ b/src/planner/operator/join.rs @@ -13,9 +13,7 @@ // limitations under the License. use super::{Operator, PlanImpl}; -use crate::expression::ScalarExpression; -use crate::iter_ext::Itertools; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt; use std::fmt::Formatter; @@ -39,9 +37,9 @@ impl JoinType { pub enum JoinCondition { On { /// Equijoin clause expressed as pairs of (left, right) join columns - on: Vec<(ScalarExpression, ScalarExpression)>, + on: Vec<(ExprRef, ExprRef)>, /// Filters applied during join (non-equi conditions) - filter: Option, + filter: Option, }, None, } @@ -83,47 +81,43 @@ impl JoinOperator { } } -impl fmt::Display for JoinType { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - match self { - JoinType::Inner => write!(f, "Inner")?, - JoinType::LeftOuter => write!(f, "LeftOuter")?, - JoinType::RightOuter => write!(f, "RightOuter")?, - JoinType::Full => write!(f, "Full")?, - JoinType::Cross => write!(f, "Cross")?, - } - - Ok(()) - } -} - -impl fmt::Display for JoinOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - write!(f, "{} Join{}", self.join_type, self.on)?; - - Ok(()) +impl Explain for JoinOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{} Join{}", self.join_type, self.on.explain(arena)) } } -impl fmt::Display for JoinCondition { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { +impl Explain for JoinCondition { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { JoinCondition::On { on, filter } => { if !on.is_empty() { - let on = on - .iter() - .map(|(v1, v2)| format!("{v1} = {v2}")) - .join(" AND "); - - write!(f, " On {on}")?; + f.write_str(" On ")?; + for (index, (left, right)) in on.iter().enumerate() { + if index > 0 { + f.write_str(" AND ")?; + } + write!(f, "{} = {}", left.explain(arena), right.explain(arena))?; + } } if let Some(filter) = filter { - write!(f, " Where {filter}")?; + write!(f, " Where {}", filter.explain(arena))?; } + Ok(()) } - JoinCondition::None => { - write!(f, " Nothing")?; - } + JoinCondition::None => f.write_str(" Nothing"), + } + } +} + +impl fmt::Display for JoinType { + fn fmt(&self, f: &mut Formatter) -> fmt::Result { + match self { + JoinType::Inner => write!(f, "Inner")?, + JoinType::LeftOuter => write!(f, "LeftOuter")?, + JoinType::RightOuter => write!(f, "RightOuter")?, + JoinType::Full => write!(f, "Full")?, + JoinType::Cross => write!(f, "Cross")?, } Ok(()) @@ -136,9 +130,14 @@ mod tests { #[test] fn forced_nested_loop_overrides_equi_join() { + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = crate::planner::PlanArena::new(&table_arena); let mut operator = JoinOperator { on: JoinCondition::On { - on: vec![(1_i32.into(), 2_i32.into())], + on: vec![( + arena.alloc_expression(crate::expression::ScalarExpression::from(1_i32)), + arena.alloc_expression(crate::expression::ScalarExpression::from(2_i32)), + )], filter: None, }, join_type: JoinType::Inner, diff --git a/src/planner/operator/mark_apply.rs b/src/planner/operator/mark_apply.rs index 68758d70..21e61ddf 100644 --- a/src/planner/operator/mark_apply.rs +++ b/src/planner/operator/mark_apply.rs @@ -14,8 +14,7 @@ use super::Operator; use crate::catalog::ColumnRef; -use crate::expression::ScalarExpression; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{Childrens, ExprRef, LogicalPlan}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt; use std::fmt::Formatter; @@ -35,13 +34,13 @@ pub enum MarkApplyKind { #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct MarkApplyOperator { pub kind: MarkApplyKind, - pub predicates: Vec, + pub predicates: Vec, output_column: ColumnRef, - pub parameterized_probe: Option, + pub parameterized_probe: Option, } impl MarkApplyOperator { - pub fn new_exists(output_column: ColumnRef, predicates: Vec) -> Self { + pub fn new_exists(output_column: ColumnRef, predicates: Vec) -> Self { Self { kind: MarkApplyKind::Exists, predicates, @@ -54,7 +53,7 @@ impl MarkApplyOperator { left: LogicalPlan, right: LogicalPlan, output_column: ColumnRef, - predicates: Vec, + predicates: Vec, ) -> LogicalPlan { LogicalPlan::new( Operator::MarkApply(MarkApplyOperator::new_exists(output_column, predicates)), @@ -65,14 +64,14 @@ impl MarkApplyOperator { ) } - pub fn new_in(output_column: ColumnRef, predicates: Vec) -> Self { + pub fn new_in(output_column: ColumnRef, predicates: Vec) -> Self { Self::new_quantified(MarkApplyQuantifier::Any, output_column, predicates) } pub fn new_quantified( quantifier: MarkApplyQuantifier, output_column: ColumnRef, - predicates: Vec, + predicates: Vec, ) -> Self { Self { kind: MarkApplyKind::Quantified(quantifier), @@ -86,7 +85,7 @@ impl MarkApplyOperator { left: LogicalPlan, right: LogicalPlan, output_column: ColumnRef, - predicates: Vec, + predicates: Vec, ) -> LogicalPlan { Self::build_quantified( left, @@ -102,7 +101,7 @@ impl MarkApplyOperator { right: LogicalPlan, quantifier: MarkApplyQuantifier, output_column: ColumnRef, - predicates: Vec, + predicates: Vec, ) -> LogicalPlan { LogicalPlan::new( Operator::MarkApply(MarkApplyOperator::new_quantified( @@ -117,11 +116,11 @@ impl MarkApplyOperator { ) } - pub fn predicates(&self) -> &[ScalarExpression] { + pub fn predicates(&self) -> &[ExprRef] { &self.predicates } - pub fn predicates_mut(&mut self) -> &mut Vec { + pub fn predicates_mut(&mut self) -> &mut Vec { &mut self.predicates } @@ -129,11 +128,11 @@ impl MarkApplyOperator { &self.output_column } - pub fn parameterized_probe(&self) -> Option<&ScalarExpression> { + pub fn parameterized_probe(&self) -> Option<&ExprRef> { self.parameterized_probe.as_ref() } - pub fn set_parameterized_probe(&mut self, probe: Option) { + pub fn set_parameterized_probe(&mut self, probe: Option) { self.parameterized_probe = probe; } } diff --git a/src/planner/operator/mod.rs b/src/planner/operator/mod.rs index 0cdaca43..68a0fbeb 100644 --- a/src/planner/operator/mod.rs +++ b/src/planner/operator/mod.rs @@ -60,7 +60,6 @@ use self::{ use crate::catalog::ColumnRef; use crate::errors::DatabaseError; use crate::expression::visitor::{walk_expr, ExprVisitor}; -use crate::expression::ScalarExpression; use crate::planner::operator::alter_table::change_column::DefaultChange as ColumnDefaultChange; use crate::planner::operator::alter_table::drop_column::DropColumnOperator; use crate::planner::operator::analyze::AnalyzeOperator; @@ -87,11 +86,9 @@ use crate::planner::operator::union::UnionOperator; use crate::planner::operator::update::UpdateOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::visitor::OperatorVisitor; -use crate::planner::{MetaArena, PlanArena}; -use crate::types::index::IndexInfo; +use crate::planner::{fmt_explain_list, Explain, ExprRef, MetaArena, PlanArena}; +use crate::types::index::{IndexInfo, IndexMetaRef}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub enum Operator { @@ -210,30 +207,131 @@ pub enum PlanImpl { Window, } +impl Explain for ColumnRef { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let column = arena.column(*self); + if let Some(table_name) = column.table_name() { + write!(f, "{}.{}", table_name, column.name()) + } else { + f.write_str(column.name()) + } + } +} + +impl Explain for IndexMetaRef { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(&arena.index(*self).name) + } +} + +macro_rules! impl_display_explain { + ($( $(#[$meta:meta])* $ty:ty),* $(,)?) => { + $( + $(#[$meta])* + impl Explain for $ty { + fn fmt(&self, _arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + std::fmt::Display::fmt(self, f) + } + } + )* + }; +} + +impl_display_explain!( + ScalarApplyOperator, + MarkApplyOperator, + ScalarSubqueryOperator, + FunctionScanOperator, + LimitOperator, + ValuesOperator, + DescribeOperator, + InsertOperator, + DeleteOperator, + AddColumnOperator, + DropColumnOperator, + CreateTableOperator, + CreateViewOperator, + DropTableOperator, + DropViewOperator, + DropIndexOperator, + TruncateOperator, + #[cfg(feature = "copy")] + CopyToFileOperator, +); + +impl Explain for Operator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Operator::Dummy => f.write_str("Dummy"), + Operator::Aggregate(op) => Explain::fmt(op, arena, f), + Operator::ScalarApply(op) => Explain::fmt(op, arena, f), + Operator::MarkApply(op) => Explain::fmt(op, arena, f), + Operator::Filter(op) => Explain::fmt(op, arena, f), + Operator::Join(op) => Explain::fmt(op, arena, f), + Operator::Project(op) => Explain::fmt(op, arena, f), + Operator::ScalarSubquery(op) => Explain::fmt(op, arena, f), + Operator::TableScan(op) => Explain::fmt(op, arena, f), + Operator::FunctionScan(op) => Explain::fmt(op, arena, f), + Operator::Sort(op) => Explain::fmt(op, arena, f), + Operator::Limit(op) => Explain::fmt(op, arena, f), + Operator::TopK(op) => Explain::fmt(op, arena, f), + Operator::Values(op) => Explain::fmt(op, arena, f), + Operator::ShowTable => f.write_str("Show Tables"), + Operator::ShowView => f.write_str("Show Views"), + Operator::Explain => unreachable!(), + Operator::Describe(op) => Explain::fmt(op, arena, f), + Operator::SetMembership(op) => Explain::fmt(op, arena, f), + Operator::Union(op) => Explain::fmt(op, arena, f), + Operator::RecursiveCte(op) => Explain::fmt(op, arena, f), + Operator::RecursiveScan(op) => Explain::fmt(op, arena, f), + Operator::Insert(op) => Explain::fmt(op, arena, f), + Operator::Update(op) => Explain::fmt(op, arena, f), + Operator::Delete(op) => Explain::fmt(op, arena, f), + Operator::Analyze(op) => Explain::fmt(op, arena, f), + Operator::AddColumn(op) => Explain::fmt(op, arena, f), + Operator::ChangeColumn(op) => Explain::fmt(op, arena, f), + Operator::DropColumn(op) => Explain::fmt(op, arena, f), + Operator::CreateTable(op) => Explain::fmt(op, arena, f), + Operator::CreateIndex(op) => Explain::fmt(op, arena, f), + Operator::CreateView(op) => Explain::fmt(op, arena, f), + Operator::DropTable(op) => Explain::fmt(op, arena, f), + Operator::DropView(op) => Explain::fmt(op, arena, f), + Operator::DropIndex(op) => Explain::fmt(op, arena, f), + Operator::Truncate(op) => Explain::fmt(op, arena, f), + #[cfg(feature = "copy")] + Operator::CopyFromFile(op) => Explain::fmt(op, arena, f), + #[cfg(feature = "copy")] + Operator::CopyToFile(op) => Explain::fmt(op, arena, f), + Operator::Window(op) => Explain::fmt(op, arena, f), + } + } +} + impl Operator { pub fn visit_referenced_columns( &self, - arena: &mut A, - f: &mut impl FnMut(&mut A, &ColumnRef) -> bool, + arena: &A, + f: &mut impl FnMut(&A, &ColumnRef) -> bool, ) -> Result { struct ReferencedColumnVisitor<'a, A, F> { - arena: &'a mut A, + arena: &'a A, f: &'a mut F, keep_going: bool, } - impl<'expr, A, F> ExprVisitor<'expr> for ReferencedColumnVisitor<'_, A, F> + impl ExprVisitor for ReferencedColumnVisitor<'_, A, F> where - F: FnMut(&mut A, &ColumnRef) -> bool, + A: MetaArena, + F: FnMut(&A, &ColumnRef) -> bool, { - fn visit(&mut self, expr: &'expr ScalarExpression) -> Result<(), DatabaseError> { + fn visit(&mut self, expr: ExprRef, arena: &A) -> Result<(), DatabaseError> { if self.keep_going { - walk_expr(self, expr)?; + walk_expr(self, expr, arena)?; } Ok(()) } - fn visit_column_ref(&mut self, column: &'expr ColumnRef) -> Result<(), DatabaseError> { + fn visit_column_ref(&mut self, column: &ColumnRef) -> Result<(), DatabaseError> { if self.keep_going { self.keep_going = (self.f)(self.arena, column); } @@ -243,14 +341,15 @@ impl Operator { impl<'operator, A, F> OperatorVisitor<'operator> for ReferencedColumnVisitor<'_, A, F> where - F: FnMut(&mut A, &ColumnRef) -> bool, + A: MetaArena, + F: FnMut(&A, &ColumnRef) -> bool, { fn visit_aggregate( &mut self, op: &'operator AggregateOperator, ) -> Result<(), DatabaseError> { for expr in op.agg_calls.iter().chain(&op.groupby_exprs) { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } @@ -260,26 +359,26 @@ impl Operator { op: &'operator MarkApplyOperator, ) -> Result<(), DatabaseError> { for expr in &op.predicates { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } if let Some(expr) = &op.parameterized_probe { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } fn visit_filter(&mut self, op: &'operator FilterOperator) -> Result<(), DatabaseError> { - ExprVisitor::visit(self, &op.predicate) + ExprVisitor::visit(self, op.predicate, self.arena) } fn visit_join(&mut self, op: &'operator JoinOperator) -> Result<(), DatabaseError> { if let JoinCondition::On { on, filter } = &op.on { for (left_expr, right_expr) in on { - ExprVisitor::visit(self, left_expr)?; - ExprVisitor::visit(self, right_expr)?; + ExprVisitor::visit(self, *left_expr, self.arena)?; + ExprVisitor::visit(self, *right_expr, self.arena)?; } if let Some(expr) = filter { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } } Ok(()) @@ -290,7 +389,7 @@ impl Operator { op: &'operator ProjectOperator, ) -> Result<(), DatabaseError> { for expr in &op.exprs { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } @@ -300,7 +399,7 @@ impl Operator { op: &'operator TableScanOperator, ) -> Result<(), DatabaseError> { for column in &op.columns { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } @@ -310,14 +409,14 @@ impl Operator { op: &'operator FunctionScanOperator, ) -> Result<(), DatabaseError> { for expr in &op.table_function.args { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } fn visit_sort(&mut self, op: &'operator SortOperator) -> Result<(), DatabaseError> { for field in &op.sort_fields { - ExprVisitor::visit(self, &field.expr)?; + ExprVisitor::visit(self, field.expr, self.arena)?; } Ok(()) } @@ -332,28 +431,28 @@ impl Operator { .map(|field| &field.expr) .chain(op.functions.iter().flat_map(|function| &function.args)) { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } fn visit_top_k(&mut self, op: &'operator TopKOperator) -> Result<(), DatabaseError> { for field in &op.sort_fields { - ExprVisitor::visit(self, &field.expr)?; + ExprVisitor::visit(self, field.expr, self.arena)?; } Ok(()) } fn visit_values(&mut self, op: &'operator ValuesOperator) -> Result<(), DatabaseError> { for column in &op.schema_ref { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } fn visit_union(&mut self, op: &'operator UnionOperator) -> Result<(), DatabaseError> { for column in op.left_schema_ref.iter().chain(&op._right_schema_ref) { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } @@ -363,7 +462,7 @@ impl Operator { op: &'operator RecursiveCteOperator, ) -> Result<(), DatabaseError> { for column in &op.schema_ref { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } @@ -373,7 +472,7 @@ impl Operator { op: &'operator RecursiveScanOperator, ) -> Result<(), DatabaseError> { for column in &op.schema_ref { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } @@ -383,21 +482,21 @@ impl Operator { op: &'operator SetMembershipOperator, ) -> Result<(), DatabaseError> { for column in op.left_schema_ref.iter().chain(&op._right_schema_ref) { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } fn visit_delete(&mut self, op: &'operator DeleteOperator) -> Result<(), DatabaseError> { for column in &op.primary_keys { - self.visit_column_ref(column)?; + ExprVisitor::visit_column_ref(self, column)?; } Ok(()) } fn visit_update(&mut self, op: &'operator UpdateOperator) -> Result<(), DatabaseError> { for (_, expr) in &op.value_exprs { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } @@ -407,7 +506,7 @@ impl Operator { op: &'operator AddColumnOperator, ) -> Result<(), DatabaseError> { if let Some(expr) = &op.column.desc().default { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } @@ -417,7 +516,7 @@ impl Operator { op: &'operator ChangeColumnOperator, ) -> Result<(), DatabaseError> { if let ColumnDefaultChange::Set(expr) = &op.default_change { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } Ok(()) } @@ -428,7 +527,7 @@ impl Operator { ) -> Result<(), DatabaseError> { for column in &op.columns { if let Some(expr) = &column.desc().default { - ExprVisitor::visit(self, expr)?; + ExprVisitor::visit(self, *expr, self.arena)?; } } Ok(()) @@ -446,7 +545,7 @@ impl Operator { pub fn any_referenced_column( &self, - arena: &mut PlanArena, + arena: &PlanArena, mut predicate: impl FnMut(&ColumnRef) -> bool, ) -> Result { let mut found = false; @@ -459,7 +558,7 @@ impl Operator { pub fn all_referenced_columns( &self, - arena: &mut PlanArena, + arena: &PlanArena, mut predicate: impl FnMut(&ColumnRef) -> bool, ) -> Result { let mut all = true; @@ -471,122 +570,73 @@ impl Operator { } } -impl fmt::Display for Operator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { +impl Explain for PlanImpl { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Operator::Dummy => write!(f, "Dummy"), - Operator::Aggregate(op) => write!(f, "{op}"), - Operator::ScalarApply(op) => write!(f, "{op}"), - Operator::MarkApply(op) => write!(f, "{op}"), - Operator::Filter(op) => write!(f, "{op}"), - Operator::Join(op) => write!(f, "{op}"), - Operator::Project(op) => write!(f, "{op}"), - Operator::ScalarSubquery(op) => write!(f, "{op}"), - Operator::TableScan(op) => write!(f, "{op}"), - Operator::FunctionScan(op) => write!(f, "{op}"), - Operator::Sort(op) => write!(f, "{op}"), - Operator::Limit(op) => write!(f, "{op}"), - Operator::TopK(op) => write!(f, "{op}"), - Operator::Values(op) => write!(f, "{op}"), - Operator::ShowTable => write!(f, "Show Tables"), - Operator::ShowView => write!(f, "Show Views"), - Operator::Explain => unreachable!(), - Operator::Describe(op) => write!(f, "{op}"), - Operator::Insert(op) => write!(f, "{op}"), - Operator::Update(op) => write!(f, "{op}"), - Operator::Delete(op) => write!(f, "{op}"), - Operator::Analyze(op) => write!(f, "{op}"), - Operator::AddColumn(op) => write!(f, "{op}"), - Operator::ChangeColumn(op) => write!(f, "{op}"), - Operator::DropColumn(op) => write!(f, "{op}"), - Operator::CreateTable(op) => write!(f, "{op}"), - Operator::CreateIndex(op) => write!(f, "{op}"), - Operator::CreateView(op) => write!(f, "{op}"), - Operator::DropTable(op) => write!(f, "{op}"), - Operator::DropView(op) => write!(f, "{op}"), - Operator::DropIndex(op) => write!(f, "{op}"), - Operator::Truncate(op) => write!(f, "{op}"), + PlanImpl::Dummy => f.write_str("Dummy"), + PlanImpl::SimpleAggregate => f.write_str("SimpleAggregate"), + PlanImpl::HashAggregate => f.write_str("HashAggregate"), + PlanImpl::StreamAggregate => f.write_str("StreamAggregate"), + PlanImpl::StreamDistinct => f.write_str("StreamDistinct"), + PlanImpl::ScalarApply => f.write_str("ScalarApply"), + PlanImpl::MarkApply => f.write_str("MarkApply"), + PlanImpl::Filter => f.write_str("Filter"), + PlanImpl::HashJoin => f.write_str("HashJoin"), + PlanImpl::NestLoopJoin => f.write_str("NestLoopJoin"), + PlanImpl::Project => f.write_str("Project"), + PlanImpl::ScalarSubquery => f.write_str("ScalarSubquery"), + PlanImpl::SeqScan => f.write_str("SeqScan"), + PlanImpl::FunctionScan => f.write_str("FunctionScan"), + PlanImpl::IndexScan(index) => write!(f, "IndexScan By {}", index.explain(arena)), + PlanImpl::Sort => f.write_str("Sort"), + PlanImpl::Limit => f.write_str("Limit"), + PlanImpl::TopK => f.write_str("TopK"), + PlanImpl::Values => f.write_str("Values"), + PlanImpl::Insert => f.write_str("Insert"), + PlanImpl::Update => f.write_str("Update"), + PlanImpl::Delete => f.write_str("Delete"), + PlanImpl::AddColumn => f.write_str("AddColumn"), + PlanImpl::ChangeColumn => f.write_str("ChangeColumn"), + PlanImpl::DropColumn => f.write_str("DropColumn"), + PlanImpl::CreateTable => f.write_str("CreateTable"), + PlanImpl::DropTable => f.write_str("DropTable"), + PlanImpl::Truncate => f.write_str("Truncate"), + PlanImpl::Show => f.write_str("Show"), #[cfg(feature = "copy")] - Operator::CopyFromFile(op) => write!(f, "{op}"), + PlanImpl::CopyFromFile => f.write_str("CopyFromFile"), #[cfg(feature = "copy")] - Operator::CopyToFile(op) => write!(f, "{op}"), - Operator::Union(op) => write!(f, "{op}"), - Operator::RecursiveCte(op) => write!(f, "{op}"), - Operator::RecursiveScan(op) => write!(f, "{op}"), - Operator::SetMembership(op) => write!(f, "{op}"), - Operator::Window(op) => write!(f, "{op}"), + PlanImpl::CopyToFile => f.write_str("CopyToFile"), + PlanImpl::Analyze => f.write_str("Analyze"), + PlanImpl::Window => f.write_str("Window"), } } } -impl fmt::Display for PhysicalOption { - fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { - write!(f, "{} => (Sort Option: {})", self.plan, self.sort_option)?; - Ok(()) - } -} - -impl fmt::Display for SortOption { - fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { +impl Explain for SortOption { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { SortOption::OrderBy { fields, ignore_prefix_len, } => { - write!(f, "OrderBy: (")?; - for (i, sort_field) in fields.iter().enumerate() { - write!(f, "{sort_field}")?; - if fields.len() - 1 != i { - write!(f, ", ")?; - } - } + f.write_str("OrderBy: (")?; + fmt_explain_list(fields, ", ", arena, f)?; write!(f, ") ignore_prefix_len: {ignore_prefix_len}") } - SortOption::Follow => write!(f, "Follow"), - SortOption::None => write!(f, "None"), + SortOption::Follow => f.write_str("Follow"), + SortOption::None => f.write_str("None"), } } } -impl fmt::Display for PlanImpl { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - match self { - PlanImpl::Dummy => write!(f, "Dummy"), - PlanImpl::SimpleAggregate => write!(f, "SimpleAggregate"), - PlanImpl::HashAggregate => write!(f, "HashAggregate"), - PlanImpl::StreamAggregate => write!(f, "StreamAggregate"), - PlanImpl::StreamDistinct => write!(f, "StreamDistinct"), - PlanImpl::ScalarApply => write!(f, "ScalarApply"), - PlanImpl::MarkApply => write!(f, "MarkApply"), - PlanImpl::Filter => write!(f, "Filter"), - PlanImpl::HashJoin => write!(f, "HashJoin"), - PlanImpl::NestLoopJoin => write!(f, "NestLoopJoin"), - PlanImpl::Project => write!(f, "Project"), - PlanImpl::ScalarSubquery => write!(f, "ScalarSubquery"), - PlanImpl::SeqScan => write!(f, "SeqScan"), - PlanImpl::FunctionScan => write!(f, "FunctionScan"), - PlanImpl::IndexScan(index) => write!(f, "IndexScan By {index}"), - PlanImpl::Sort => write!(f, "Sort"), - PlanImpl::Limit => write!(f, "Limit"), - PlanImpl::TopK => write!(f, "TopK"), - PlanImpl::Values => write!(f, "Values"), - PlanImpl::Insert => write!(f, "Insert"), - PlanImpl::Update => write!(f, "Update"), - PlanImpl::Delete => write!(f, "Delete"), - PlanImpl::AddColumn => write!(f, "AddColumn"), - PlanImpl::ChangeColumn => write!(f, "ChangeColumn"), - PlanImpl::DropColumn => write!(f, "DropColumn"), - PlanImpl::CreateTable => write!(f, "CreateTable"), - PlanImpl::DropTable => write!(f, "DropTable"), - PlanImpl::Truncate => write!(f, "Truncate"), - PlanImpl::Show => write!(f, "Show"), - #[cfg(feature = "copy")] - PlanImpl::CopyFromFile => write!(f, "CopyFromFile"), - #[cfg(feature = "copy")] - PlanImpl::CopyToFile => write!(f, "CopyToFile"), - PlanImpl::Analyze => write!(f, "Analyze"), - PlanImpl::Window => write!(f, "Window"), - } +impl Explain for PhysicalOption { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "{} => (Sort Option: {})", + self.plan.explain(arena), + self.sort_option.explain(arena) + ) } } @@ -607,8 +657,9 @@ mod tests { use crate::planner::operator::set_membership::SetMembershipKind; use crate::planner::operator::sort::SortField; use crate::planner::operator::values::ValuesOperator; + use crate::planner::ExprRef; use crate::planner::{Childrens, LogicalPlan, TableArenaCell}; - use crate::types::index::{IndexInfo, IndexMetaRef, IndexType}; + use crate::types::index::{IndexInfo, IndexMeta, IndexMetaRef, IndexType}; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -624,9 +675,9 @@ mod tests { arena.alloc_column(column_catalog(name)) } - fn index_info() -> IndexInfo { + fn index_info(meta: IndexMetaRef) -> IndexInfo { IndexInfo { - meta: IndexMetaRef::new(4), + meta, sort_option: SortOption::None, lookup: None, residual_predicate: None, @@ -637,8 +688,8 @@ mod tests { } } - fn column_expr(column: ColumnRef, position: usize) -> ScalarExpression { - ScalarExpression::column_expr(column, position) + fn column_expr(column: ColumnRef, position: usize, arena: &mut PlanArena) -> ExprRef { + arena.alloc_expression(ScalarExpression::column_expr(column, position)) } fn referenced_columns( @@ -654,29 +705,36 @@ mod tests { } #[test] - fn physical_option_and_sort_option_display() { - let sort_field = SortField::new(ScalarExpression::from(1i32), false, true); + fn physical_option_and_sort_option_explain() { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let sort_option = SortOption::OrderBy { - fields: vec![sort_field], + fields: vec![SortField::new( + arena.alloc_expression(ScalarExpression::from(1i32)), + false, + true, + )], ignore_prefix_len: 2, }; assert_eq!( - sort_option.to_string(), + sort_option.explain(&arena).to_string(), "OrderBy: (1 Desc Nulls First) ignore_prefix_len: 2" ); - assert_eq!(SortOption::Follow.to_string(), "Follow"); - assert_eq!(SortOption::None.to_string(), "None"); + assert_eq!(SortOption::Follow.explain(&arena).to_string(), "Follow"); + assert_eq!(SortOption::None.explain(&arena).to_string(), "None"); let physical = PhysicalOption::new(PlanImpl::TopK, sort_option.clone()); assert_eq!( - physical.to_string(), + physical.explain(&arena).to_string(), "TopK => (Sort Option: OrderBy: (1 Desc Nulls First) ignore_prefix_len: 2)" ); assert_eq!(physical.sort_option(), &sort_option); } #[test] - fn plan_impl_display_covers_physical_variants() { + fn plan_impl_explain_covers_physical_variants() { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let cases = [ (PlanImpl::Dummy, "Dummy"), (PlanImpl::SimpleAggregate, "SimpleAggregate"), @@ -710,16 +768,34 @@ mod tests { ]; for (plan, expected) in cases { - assert_eq!(plan.to_string(), expected); + assert_eq!(plan.explain(&arena).to_string(), expected); } + + let meta = arena.alloc_index(IndexMeta { + id: 1, + column_ids: vec![1], + table_name: "users".into(), + pk_ty: LogicalType::Integer, + value_ty: LogicalType::Integer, + name: "idx_users_id".to_string(), + ty: IndexType::Normal, + }); assert_eq!( - PlanImpl::IndexScan(Box::new(index_info())).to_string(), - "IndexScan By #4 => EMPTY" + PlanImpl::IndexScan(Box::new(index_info(meta))) + .explain(&arena) + .to_string(), + "IndexScan By idx_users_id => EMPTY" ); #[cfg(feature = "copy")] { - assert_eq!(PlanImpl::CopyFromFile.to_string(), "CopyFromFile"); - assert_eq!(PlanImpl::CopyToFile.to_string(), "CopyToFile"); + assert_eq!( + PlanImpl::CopyFromFile.explain(&arena).to_string(), + "CopyFromFile" + ); + assert_eq!( + PlanImpl::CopyToFile.explain(&arena).to_string(), + "CopyToFile" + ); } } @@ -734,21 +810,21 @@ mod tests { schema_ref: vec![left, right], }); - assert!(values.any_referenced_column(&mut arena, |column| *column == right)?); - assert!(!values - .any_referenced_column(&mut arena, |column| { *column != left && *column != right })?); - assert!(values.all_referenced_columns(&mut arena, |column| { - *column == left || *column == right - })?); - assert!(!values.all_referenced_columns(&mut arena, |column| *column == left)?); + assert!(values.any_referenced_column(&arena, |column| *column == right)?); + assert!( + !values.any_referenced_column(&arena, |column| *column != left && *column != right)? + ); + assert!(values + .all_referenced_columns(&arena, |column| { *column == left || *column == right })?); + assert!(!values.all_referenced_columns(&arena, |column| *column == left)?); let delete = Operator::Delete(DeleteOperator { table_name: "users".into(), primary_keys: vec![left], }); - assert!(delete.any_referenced_column(&mut arena, |column| *column == left)?); - assert!(Operator::Dummy.all_referenced_columns(&mut arena, |_| false)?); - assert!(!Operator::Dummy.any_referenced_column(&mut arena, |_| true)?); + assert!(delete.any_referenced_column(&arena, |column| *column == left)?); + assert!(Operator::Dummy.all_referenced_columns(&arena, |_| false)?); + assert!(!Operator::Dummy.any_referenced_column(&arena, |_| true)?); Ok(()) } @@ -762,22 +838,22 @@ mod tests { let d = column("d", &mut arena); let aggregate = Operator::Aggregate(AggregateOperator { - agg_calls: vec![column_expr(a, 0)], - groupby_exprs: vec![column_expr(b, 1)], + agg_calls: vec![column_expr(a, 0, &mut arena)], + groupby_exprs: vec![column_expr(b, 1, &mut arena)], is_distinct: false, force_spill: false, }); assert_eq!(referenced_columns(&aggregate, &mut arena)?, vec![a, b]); - let mut mark_apply = MarkApplyOperator::new_exists(d, vec![column_expr(c, 2)]); - mark_apply.set_parameterized_probe(Some(column_expr(d, 3))); + let mut mark_apply = MarkApplyOperator::new_exists(d, vec![column_expr(c, 2, &mut arena)]); + mark_apply.set_parameterized_probe(Some(column_expr(d, 3, &mut arena))); assert_eq!( referenced_columns(&Operator::MarkApply(mark_apply), &mut arena)?, vec![c, d] ); let filter = Operator::Filter(FilterOperator { - predicate: column_expr(a, 0), + predicate: column_expr(a, 0, &mut arena), is_optimized: false, having: false, }); @@ -787,21 +863,21 @@ mod tests { join_type: join::JoinType::Inner, force_nested_loop: false, on: JoinCondition::On { - on: vec![(column_expr(a, 0), column_expr(b, 1))], - filter: Some(column_expr(c, 2)), + on: vec![(column_expr(a, 0, &mut arena), column_expr(b, 1, &mut arena))], + filter: Some(column_expr(c, 2, &mut arena)), }, }); assert_eq!(referenced_columns(&join, &mut arena)?, vec![a, b, c]); - assert!(!join.all_referenced_columns(&mut arena, |column| *column == a)?); + assert!(!join.all_referenced_columns(&arena, |column| *column == a)?); let project = Operator::Project(ProjectOperator { - exprs: vec![column_expr(b, 1), column_expr(c, 2)], + exprs: vec![column_expr(b, 1, &mut arena), column_expr(c, 2, &mut arena)], }); assert_eq!(referenced_columns(&project, &mut arena)?, vec![b, c]); let update = Operator::Update(UpdateOperator { table_name: "users".into(), - value_exprs: vec![(b, column_expr(a, 0))], + value_exprs: vec![(b, column_expr(a, 0, &mut arena))], }); assert_eq!(referenced_columns(&update, &mut arena)?, vec![a]); @@ -815,7 +891,7 @@ mod tests { LogicalType::Integer, None, false, - Some(ScalarExpression::from(1_i32)), + Some(arena.alloc_expression(ScalarExpression::from(1_i32))), )?, ), }); @@ -826,7 +902,7 @@ mod tests { old_column_name: "old".to_string(), new_column_name: "new".to_string(), data_type: LogicalType::Integer, - default_change: DefaultChange::Set(column_expr(b, 1)), + default_change: DefaultChange::Set(column_expr(b, 1, &mut arena)), not_null_change: NotNullChange::NoChange, }); assert_eq!(referenced_columns(&change_column, &mut arena)?, vec![b]); @@ -840,7 +916,7 @@ mod tests { LogicalType::Integer, None, false, - Some(ScalarExpression::from(2_i32)), + Some(arena.alloc_expression(ScalarExpression::from(2_i32))), )?, )], if_not_exists: false, @@ -858,7 +934,7 @@ mod tests { let function_scan = Operator::FunctionScan(FunctionScanOperator { table_function: TableFunction { - args: vec![column_expr(c, 2)], + args: vec![column_expr(c, 2, &mut arena)], catalog: TableFunctionCatalog { schema: Vec::new(), inner: ArcTableFunctionImpl(Numbers::new()), @@ -868,12 +944,12 @@ mod tests { assert_eq!(referenced_columns(&function_scan, &mut arena)?, vec![c]); let sort = Operator::Sort(SortOperator { - sort_fields: vec![SortField::from(column_expr(a, 0))], + sort_fields: vec![SortField::from(column_expr(a, 0, &mut arena))], }); assert_eq!(referenced_columns(&sort, &mut arena)?, vec![a]); let top_k = Operator::TopK(TopKOperator { - sort_fields: vec![SortField::from(column_expr(b, 1))], + sort_fields: vec![SortField::from(column_expr(b, 1, &mut arena))], limit: 3, offset: None, }); @@ -914,7 +990,7 @@ mod tests { } #[test] - fn recursive_operators_visit_display_and_build() -> Result<(), DatabaseError> { + fn recursive_operators_visit_explain_and_build() -> Result<(), DatabaseError> { let table_arena = TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let value = column("value", &mut arena); @@ -925,13 +1001,19 @@ mod tests { schema_ref: schema.clone(), }); assert_eq!(referenced_columns(&recursive_cte, &mut arena)?, schema); - assert_eq!(recursive_cte.to_string(), "Recursive CTE: [#0, #1]"); + assert_eq!( + recursive_cte.explain(&arena).to_string(), + "Recursive CTE: [value, depth]" + ); let recursive_scan = Operator::RecursiveScan(RecursiveScanOperator { schema_ref: schema.clone(), }); assert_eq!(referenced_columns(&recursive_scan, &mut arena)?, schema); - assert_eq!(recursive_scan.to_string(), "Recursive Scan: [#0, #1]"); + assert_eq!( + recursive_scan.explain(&arena).to_string(), + "Recursive Scan: [value, depth]" + ); let anchor = LogicalPlan::new(Operator::ShowTable, Childrens::None); let recursive = LogicalPlan::new(Operator::ShowView, Childrens::None); @@ -943,19 +1025,23 @@ mod tests { #[test] fn mark_apply_constructors_and_accessors_cover_quantified_paths() { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); let left = LogicalPlan::new(Operator::ShowTable, Childrens::None); let right = LogicalPlan::new(Operator::ShowView, Childrens::None); let output = ColumnRef::new(10); - let probe = ScalarExpression::from(true); + let probe = arena.alloc_expression(ScalarExpression::from(true)); + let one = arena.alloc_expression(ScalarExpression::from(1_i32)); - let mut any = MarkApplyOperator::new_in(output, vec![ScalarExpression::from(1_i32)]); + let mut any = MarkApplyOperator::new_in(output, vec![one]); assert_eq!(any.to_string(), "MarkAnyApply"); assert_eq!(any.predicates().len(), 1); - any.predicates_mut().push(ScalarExpression::from(2_i32)); + any.predicates_mut() + .push(arena.alloc_expression(ScalarExpression::from(2_i32))); assert_eq!(any.predicates().len(), 2); assert_eq!(*any.output_column(), output); assert!(any.parameterized_probe().is_none()); - any.set_parameterized_probe(Some(probe.clone())); + any.set_parameterized_probe(Some(probe)); assert_eq!(any.parameterized_probe(), Some(&probe)); any.set_parameterized_probe(None); assert!(any.parameterized_probe().is_none()); @@ -963,17 +1049,12 @@ mod tests { let all = MarkApplyOperator::new_quantified( MarkApplyQuantifier::All, output, - vec![ScalarExpression::from(false)], + vec![arena.alloc_expression(ScalarExpression::from(false))], ); assert_eq!(all.to_string(), "MarkAllApply"); - let in_plan = MarkApplyOperator::build_in( - left.clone(), - right.clone(), - output, - vec![ScalarExpression::from(1_i32)], - ); - assert_eq!(in_plan.operator.to_string(), "MarkAnyApply"); + let in_plan = MarkApplyOperator::build_in(left.clone(), right.clone(), output, vec![one]); + assert_eq!(in_plan.operator.explain(&arena).to_string(), "MarkAnyApply"); assert!(matches!(*in_plan.childrens, Childrens::Twins { .. })); let all_plan = MarkApplyOperator::build_quantified( @@ -981,14 +1062,17 @@ mod tests { right, MarkApplyQuantifier::All, output, - vec![ScalarExpression::from(1_i32)], + vec![one], + ); + assert_eq!( + all_plan.operator.explain(&arena).to_string(), + "MarkAllApply" ); - assert_eq!(all_plan.operator.to_string(), "MarkAllApply"); assert!(matches!(*all_plan.childrens, Childrens::Twins { .. })); } #[test] - fn ddl_operator_display_formats_table_index_and_column_actions() { + fn ddl_operator_explain_formats_table_index_and_column_actions() { let table_arena = TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let id = column("id", &mut arena); @@ -1017,7 +1101,7 @@ mod tests { if_not_exists: false, ty: IndexType::Normal, }), - "Create Index On users -> [#0, #1], If Not Exists: false", + "Create Index On users -> [id, name], If Not Exists: false", ), ( Operator::CreateView(CreateViewOperator { @@ -1076,15 +1160,24 @@ mod tests { ]; for (operator, expected) in cases { - assert_eq!(operator.to_string(), expected); + assert_eq!(operator.explain(&arena).to_string(), expected); } } #[test] - fn dml_values_describe_and_analyze_display_formats_payloads() { + fn dml_values_describe_and_analyze_explain_formats_payloads() { let table_arena = TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let id = column("id", &mut arena); + let index = arena.alloc_index(IndexMeta { + id: 1, + column_ids: vec![1], + table_name: "users".into(), + pk_ty: LogicalType::Integer, + value_ty: LogicalType::Integer, + name: "idx_users_id".to_string(), + ty: IndexType::Normal, + }); let cases = [ ( @@ -1098,9 +1191,9 @@ mod tests { ( Operator::Update(UpdateOperator { table_name: "users".into(), - value_exprs: vec![(id, ScalarExpression::from(7_i32))], + value_exprs: vec![(id, arena.alloc_expression(ScalarExpression::from(7_i32)))], }), - "Update users set #0 -> 7", + "Update users set id -> 7", ), ( Operator::Delete(DeleteOperator { @@ -1128,39 +1221,46 @@ mod tests { ( Operator::Analyze(AnalyzeOperator { table_name: "users".into(), - index_metas: vec![IndexMetaRef::new(3)], + index_metas: vec![index], histogram_buckets: Some(128), }), - "Analyze users -> [#3]", + "Analyze users -> [idx_users_id]", ), ]; for (operator, expected) in cases { - assert_eq!(operator.to_string(), expected); + assert_eq!(operator.explain(&arena).to_string(), expected); } } #[test] - fn sort_and_top_k_display_fields_and_build_single_child_plan() { - let descending_nulls_first = SortField::from(ScalarExpression::from(9_i32)) - .desc() - .nulls_first(); - let ascending_nulls_last = SortField::new(ScalarExpression::from(1_i32), false, true) - .asc() - .nulls_last(); + fn sort_and_top_k_explain_fields_and_build_single_child_plan() { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let descending_nulls_first = + SortField::from(arena.alloc_expression(ScalarExpression::from(9_i32))) + .desc() + .nulls_first(); + let ascending_nulls_last = SortField::new( + arena.alloc_expression(ScalarExpression::from(1_i32)), + false, + true, + ) + .asc() + .nulls_last(); let sort = Operator::Sort(SortOperator { sort_fields: vec![descending_nulls_first.clone(), ascending_nulls_last.clone()], }); assert_eq!( - sort.to_string(), + sort.explain(&arena).to_string(), "Sort By 9 Desc Nulls First, 1 Asc Nulls Last" ); let child = LogicalPlan::new(Operator::ShowTable, Childrens::None); let top_k = TopKOperator::build(vec![descending_nulls_first], 5, Some(2), child); assert_eq!( - top_k.operator.to_string(), + top_k.operator.explain(&arena).to_string(), "Top 5, Offset 2, Sort By 9 Desc Nulls First" ); assert!(matches!(*top_k.childrens, Childrens::Only(_))); @@ -1171,13 +1271,15 @@ mod tests { offset: None, }); assert_eq!( - top_k_without_offset.to_string(), + top_k_without_offset.explain(&arena).to_string(), "Top 3, Sort By 1 Asc Nulls Last" ); } #[test] fn drop_index_build_preserves_operator_payload_and_children() { + let table_arena = TableArenaCell::default(); + let arena = PlanArena::new(&table_arena); let plan = DropIndexOperator::build( "users".into(), "idx_users_id".to_string(), @@ -1186,20 +1288,21 @@ mod tests { ); assert_eq!( - plan.operator.to_string(), + plan.operator.explain(&arena).to_string(), "Drop Index idx_users_id On users, If Exists: true" ); assert!(matches!(*plan.childrens, Childrens::None)); } #[test] - fn function_scan_display_and_build_preserve_table_function() { + fn function_scan_explain_and_build_preserve_table_function() { let table_arena = TableArenaCell::default(); let numbers = Numbers::new(); let mut schema = Vec::new(); numbers.output_schema_into(table_arena.borrow_mut(), &mut schema); + let mut arena = PlanArena::new(&table_arena); let table_function = TableFunction { - args: vec![ScalarExpression::from(3_i32)], + args: vec![arena.alloc_expression(ScalarExpression::from(3_i32))], catalog: TableFunctionCatalog { schema, inner: ArcTableFunctionImpl(numbers), @@ -1208,12 +1311,15 @@ mod tests { let plan = FunctionScanOperator::build(table_function); - assert_eq!(plan.operator.to_string(), "Function Scan: numbers"); + assert_eq!( + plan.operator.explain(&arena).to_string(), + "Function Scan: numbers" + ); assert!(matches!(*plan.childrens, Childrens::None)); } #[test] - fn set_membership_display_and_build_cover_both_kinds() { + fn set_membership_explain_and_build_cover_both_kinds() { let table_arena = TableArenaCell::default(); let mut arena = PlanArena::new(&table_arena); let left_col = column("left_id", &mut arena); @@ -1229,7 +1335,10 @@ mod tests { right, ); - assert_eq!(plan.operator.to_string(), "Intersect: [#0]"); + assert_eq!( + plan.operator.explain(&arena).to_string(), + "Intersect: [left_id]" + ); assert!(matches!(*plan.childrens, Childrens::Twins { .. })); assert_eq!( Operator::SetMembership(SetMembershipOperator { @@ -1237,28 +1346,34 @@ mod tests { left_schema_ref: vec![left_col], _right_schema_ref: vec![right_col], }) + .explain(&arena) .to_string(), - "Except: [#0]" + "Except: [left_id]" ); } #[test] fn scalar_apply_and_subquery_build_expected_child_shapes() { + let table_arena = TableArenaCell::default(); + let arena = PlanArena::new(&table_arena); let left = LogicalPlan::new(Operator::ShowTable, Childrens::None); let right = LogicalPlan::new(Operator::ShowView, Childrens::None); let apply = ScalarApplyOperator::build(left.clone(), right); - assert_eq!(apply.operator.to_string(), "ScalarApply"); + assert_eq!(apply.operator.explain(&arena).to_string(), "ScalarApply"); assert!(matches!(*apply.childrens, Childrens::Twins { .. })); let subquery = ScalarSubqueryOperator::build(left); - assert_eq!(subquery.operator.to_string(), "ScalarSubquery"); + assert_eq!( + subquery.operator.explain(&arena).to_string(), + "ScalarSubquery" + ); assert!(matches!(*subquery.childrens, Childrens::Only(_))); } #[cfg(feature = "copy")] #[test] - fn copy_display_formats_source_target_table_and_schema() { + fn copy_explain_formats_source_target_table_and_schema() { use crate::binder::copy::{ExtSource, FileFormat}; use std::path::PathBuf; @@ -1282,8 +1397,8 @@ mod tests { }); assert_eq!( - operator.to_string(), - "Copy /tmp/users.csv -> users [#0, #1]" + operator.explain(&arena).to_string(), + "Copy /tmp/users.csv -> users [id, name]" ); assert_eq!( Operator::CopyToFile(CopyToFileOperator { @@ -1297,6 +1412,7 @@ mod tests { }, }, }) + .explain(&arena) .to_string(), "Copy To /tmp/output.csv" ); diff --git a/src/planner/operator/project.rs b/src/planner/operator/project.rs index a1af3b05..97003019 100644 --- a/src/planner/operator/project.rs +++ b/src/planner/operator/project.rs @@ -12,23 +12,18 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::expression::ScalarExpression; -use crate::iter_ext::Itertools; +use crate::planner::{fmt_explain_list, Explain, ExprRef, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct ProjectOperator { - pub exprs: Vec, + pub exprs: Vec, } -impl fmt::Display for ProjectOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let exprs = self.exprs.iter().map(|expr| format!("{expr}")).join(", "); - - write!(f, "Projection [{exprs}]")?; - - Ok(()) +impl Explain for ProjectOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Projection [")?; + fmt_explain_list(&self.exprs, ", ", arena, f)?; + f.write_str("]") } } diff --git a/src/planner/operator/recursive_cte.rs b/src/planner/operator/recursive_cte.rs index fa0959e5..a84d3b34 100644 --- a/src/planner/operator/recursive_cte.rs +++ b/src/planner/operator/recursive_cte.rs @@ -12,12 +12,10 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::iter_ext::Itertools; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct RecursiveCteOperator { @@ -36,9 +34,11 @@ impl RecursiveCteOperator { } } -impl fmt::Display for RecursiveCteOperator { - fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { - write!(f, "Recursive CTE: [{}]", self.schema_ref.iter().join(", ")) +impl Explain for RecursiveCteOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Recursive CTE: [")?; + fmt_explain_list(&self.schema_ref, ", ", arena, f)?; + f.write_str("]") } } @@ -47,8 +47,10 @@ pub struct RecursiveScanOperator { pub schema_ref: Schema, } -impl fmt::Display for RecursiveScanOperator { - fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { - write!(f, "Recursive Scan: [{}]", self.schema_ref.iter().join(", ")) +impl Explain for RecursiveScanOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Recursive Scan: [")?; + fmt_explain_list(&self.schema_ref, ", ", arena, f)?; + f.write_str("]") } } diff --git a/src/planner/operator/set_membership.rs b/src/planner/operator/set_membership.rs index 34ad6fb7..e659dc6e 100644 --- a/src/planner/operator/set_membership.rs +++ b/src/planner/operator/set_membership.rs @@ -12,13 +12,10 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::iter_ext::Itertools; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Copy, Clone, Hash, ReferenceSerialization)] pub enum SetMembershipKind { @@ -65,12 +62,10 @@ impl SetMembershipOperator { } } -impl fmt::Display for SetMembershipOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let schema = self.left_schema_ref.iter().join(", "); - - write!(f, "{}: [{schema}]", self.kind.name())?; - - Ok(()) +impl Explain for SetMembershipOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}: [", self.kind.name())?; + fmt_explain_list(&self.left_schema_ref, ", ", arena, f)?; + f.write_str("]") } } diff --git a/src/planner/operator/sort.rs b/src/planner/operator/sort.rs index a680591a..3de06a32 100644 --- a/src/planner/operator/sort.rs +++ b/src/planner/operator/sort.rs @@ -12,21 +12,18 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::expression::ScalarExpression; -use crate::iter_ext::Itertools; +use crate::planner::{fmt_explain_list, Explain, ExprRef, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct SortField { - pub expr: ScalarExpression, + pub expr: ExprRef, pub asc: bool, pub nulls_first: bool, } impl SortField { - pub fn new(expr: ScalarExpression, asc: bool, nulls_first: bool) -> Self { + pub fn new(expr: ExprRef, asc: bool, nulls_first: bool) -> Self { SortField { expr, asc, @@ -55,8 +52,8 @@ impl SortField { } } -impl From for SortField { - fn from(expr: ScalarExpression) -> Self { +impl From for SortField { + fn from(expr: ExprRef) -> Self { SortField::new(expr, true, false) } } @@ -66,31 +63,21 @@ pub struct SortOperator { pub sort_fields: Vec, } -impl fmt::Display for SortOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let sort_fields = self - .sort_fields - .iter() - .map(|sort_field| format!("{sort_field}")) - .join(", "); - write!(f, "Sort By {sort_fields}") +impl Explain for SortOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Sort By ")?; + fmt_explain_list(&self.sort_fields, ", ", arena, f) } } -impl fmt::Display for SortField { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - write!(f, "{}", self.expr)?; - if self.asc { - write!(f, " Asc")?; +impl Explain for SortField { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let direction = if self.asc { "Asc" } else { "Desc" }; + let nulls = if self.nulls_first { + "Nulls First" } else { - write!(f, " Desc")?; - } - if self.nulls_first { - write!(f, " Nulls First")?; - } else { - write!(f, " Nulls Last")?; - } - - Ok(()) + "Nulls Last" + }; + write!(f, "{} {direction} {nulls}", self.expr.explain(arena)) } } diff --git a/src/planner/operator/table_scan.rs b/src/planner/operator/table_scan.rs index 5ad7e552..d5d8c397 100644 --- a/src/planner/operator/table_scan.rs +++ b/src/planner/operator/table_scan.rs @@ -18,12 +18,10 @@ use crate::errors::DatabaseError; use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; use crate::planner::operator::sort::SortField; -use crate::planner::{Childrens, LogicalPlan, PlanArena}; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; use crate::storage::Bounds; use crate::types::index::IndexInfo; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct TableScanOperator { @@ -43,21 +41,22 @@ impl TableScanOperator { table_name: TableName, table_catalog: &TableCatalog, with_pk: bool, - arena: &PlanArena, + arena: &mut PlanArena, ) -> Result { // Fill all Columns in TableCatalog by default let columns = table_catalog.columns().copied().collect_vec(); let mut index_infos = Vec::with_capacity(table_catalog.indexes.len()); for index_ref in table_catalog.indexes.iter().copied() { - let index_meta = arena.index(index_ref); - let mut sort_fields = Vec::with_capacity(index_meta.column_ids.len()); - for col_id in &index_meta.column_ids { + let column_ids = arena.index(index_ref).column_ids.clone(); + let mut sort_fields = Vec::with_capacity(column_ids.len()); + for (position, col_id) in column_ids.iter().enumerate() { let column_ref = table_catalog.get_column_by_id(col_id).ok_or_else(|| { DatabaseError::column_not_found(format!("index column id: {col_id} not found")) })?; sort_fields.push(SortField { - expr: ScalarExpression::column_expr(column_ref, sort_fields.len()), + expr: arena + .alloc_expression(ScalarExpression::column_expr(column_ref, position)), asc: true, nulls_first: false, }) @@ -91,23 +90,18 @@ impl TableScanOperator { } } -impl fmt::Display for TableScanOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let projection_columns = self.columns.iter().join(", "); +impl Explain for TableScanOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "TableScan {} -> [", self.table_name)?; + fmt_explain_list(&self.columns, ", ", arena, f)?; + f.write_str("]")?; let (offset, limit) = self.limit; - - write!( - f, - "TableScan {} -> [{}]", - self.table_name, projection_columns - )?; if let Some(limit) = limit { write!(f, ", Limit: {limit}")?; } if let Some(offset) = offset { write!(f, ", Offset: {offset}")?; } - Ok(()) } } diff --git a/src/planner/operator/top_k.rs b/src/planner/operator/top_k.rs index 47364cae..39080631 100644 --- a/src/planner/operator/top_k.rs +++ b/src/planner/operator/top_k.rs @@ -13,12 +13,9 @@ // limitations under the License. use super::Operator; -use crate::iter_ext::Itertools; use crate::planner::operator::sort::SortField; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct TopKOperator { @@ -45,21 +42,13 @@ impl TopKOperator { } } -impl fmt::Display for TopKOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { +impl Explain for TopKOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "Top {}, ", self.limit)?; - if let Some(offset) = self.offset { write!(f, "Offset {offset}, ")?; } - - let sort_fields = self - .sort_fields - .iter() - .map(|sort_field| format!("{sort_field}")) - .join(", "); - write!(f, "Sort By {sort_fields}")?; - - Ok(()) + f.write_str("Sort By ")?; + fmt_explain_list(&self.sort_fields, ", ", arena, f) } } diff --git a/src/planner/operator/union.rs b/src/planner/operator/union.rs index 748fe13f..30370f55 100644 --- a/src/planner/operator/union.rs +++ b/src/planner/operator/union.rs @@ -12,13 +12,10 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::iter_ext::Itertools; use crate::planner::operator::Operator; -use crate::planner::{Childrens, LogicalPlan}; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct UnionOperator { @@ -47,12 +44,10 @@ impl UnionOperator { } } -impl fmt::Display for UnionOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let schema = self.left_schema_ref.iter().join(", "); - - write!(f, "Union: [{schema}]")?; - - Ok(()) +impl Explain for UnionOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("Union: [")?; + fmt_explain_list(&self.left_schema_ref, ", ", arena, f)?; + f.write_str("]") } } diff --git a/src/planner/operator/update.rs b/src/planner/operator/update.rs index 4f668b73..e4c2204e 100644 --- a/src/planner/operator/update.rs +++ b/src/planner/operator/update.rs @@ -13,27 +13,24 @@ // limitations under the License. use crate::catalog::{ColumnRef, TableName}; -use crate::expression::ScalarExpression; -use crate::iter_ext::Itertools; +use crate::planner::{Explain, ExprRef, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct UpdateOperator { pub table_name: TableName, - pub value_exprs: Vec<(ColumnRef, ScalarExpression)>, + pub value_exprs: Vec<(ColumnRef, ExprRef)>, } -impl fmt::Display for UpdateOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let values = self - .value_exprs - .iter() - .map(|(column, expr)| format!("{column} -> {expr}")) - .join(", "); - write!(f, "Update {} set {}", self.table_name, values)?; - +impl Explain for UpdateOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Update {} set ", self.table_name)?; + for (index, (column, expr)) in self.value_exprs.iter().enumerate() { + if index > 0 { + f.write_str(", ")?; + } + write!(f, "{} -> {}", column.explain(arena), expr.explain(arena))?; + } Ok(()) } } diff --git a/src/planner/operator/visitor.rs b/src/planner/operator/visitor.rs index a1da1aa9..c81f075f 100644 --- a/src/planner/operator/visitor.rs +++ b/src/planner/operator/visitor.rs @@ -16,6 +16,7 @@ use super::alter_table::change_column::DefaultChange; use super::*; use crate::errors::DatabaseError; use crate::expression::visitor::ExprVisitor; +use crate::planner::MetaArena; pub trait OperatorVisitor<'a>: Sized { fn visit_operator(&mut self, operator: &'a Operator) -> Result<(), DatabaseError> { @@ -190,46 +191,47 @@ pub trait OperatorVisitor<'a>: Sized { } } -pub struct OperatorExprVisitor<'a, V> { +pub struct OperatorExprVisitor<'a, V, A> { visitor: &'a mut V, + arena: &'a A, } -impl<'a, V> OperatorExprVisitor<'a, V> { - pub fn new(visitor: &'a mut V) -> Self { - Self { visitor } +impl<'a, V, A> OperatorExprVisitor<'a, V, A> { + pub fn new(visitor: &'a mut V, arena: &'a A) -> Self { + Self { visitor, arena } } } -impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> { +impl<'a, V: ExprVisitor, A: MetaArena> OperatorVisitor<'a> for OperatorExprVisitor<'_, V, A> { fn visit_aggregate(&mut self, op: &'a AggregateOperator) -> Result<(), DatabaseError> { for expr in op.agg_calls.iter().chain(&op.groupby_exprs) { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_mark_apply(&mut self, op: &'a MarkApplyOperator) -> Result<(), DatabaseError> { for expr in &op.predicates { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } if let Some(expr) = &op.parameterized_probe { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_filter(&mut self, op: &'a FilterOperator) -> Result<(), DatabaseError> { - ExprVisitor::visit(self.visitor, &op.predicate) + ExprVisitor::visit(self.visitor, op.predicate, self.arena) } fn visit_join(&mut self, op: &'a JoinOperator) -> Result<(), DatabaseError> { if let JoinCondition::On { on, filter } = &op.on { for (left_expr, right_expr) in on { - ExprVisitor::visit(self.visitor, left_expr)?; - ExprVisitor::visit(self.visitor, right_expr)?; + ExprVisitor::visit(self.visitor, *left_expr, self.arena)?; + ExprVisitor::visit(self.visitor, *right_expr, self.arena)?; } if let Some(expr) = filter { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } } Ok(()) @@ -237,7 +239,7 @@ impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> fn visit_project(&mut self, op: &'a ProjectOperator) -> Result<(), DatabaseError> { for expr in &op.exprs { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } @@ -246,11 +248,11 @@ impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> for index_info in &op.index_infos { if let SortOption::OrderBy { fields, .. } = &index_info.sort_option { for field in fields { - ExprVisitor::visit(self.visitor, &field.expr)?; + ExprVisitor::visit(self.visitor, field.expr, self.arena)?; } } if let Some(expr) = &index_info.residual_predicate { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } } Ok(()) @@ -258,14 +260,14 @@ impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> fn visit_function_scan(&mut self, op: &'a FunctionScanOperator) -> Result<(), DatabaseError> { for expr in &op.table_function.args { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_sort(&mut self, op: &'a SortOperator) -> Result<(), DatabaseError> { for field in &op.sort_fields { - ExprVisitor::visit(self.visitor, &field.expr)?; + ExprVisitor::visit(self.visitor, field.expr, self.arena)?; } Ok(()) } @@ -277,35 +279,35 @@ impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> .map(|field| &field.expr) .chain(op.functions.iter().flat_map(|function| &function.args)) { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_top_k(&mut self, op: &'a TopKOperator) -> Result<(), DatabaseError> { for field in &op.sort_fields { - ExprVisitor::visit(self.visitor, &field.expr)?; + ExprVisitor::visit(self.visitor, field.expr, self.arena)?; } Ok(()) } fn visit_update(&mut self, op: &'a UpdateOperator) -> Result<(), DatabaseError> { for (_, expr) in &op.value_exprs { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_add_column(&mut self, op: &'a AddColumnOperator) -> Result<(), DatabaseError> { if let Some(expr) = &op.column.desc().default { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } fn visit_change_column(&mut self, op: &'a ChangeColumnOperator) -> Result<(), DatabaseError> { if let DefaultChange::Set(expr) = &op.default_change { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } Ok(()) } @@ -313,7 +315,7 @@ impl<'a, V: ExprVisitor<'a>> OperatorVisitor<'a> for OperatorExprVisitor<'_, V> fn visit_create_table(&mut self, op: &'a CreateTableOperator) -> Result<(), DatabaseError> { for column in &op.columns { if let Some(expr) = &column.desc().default { - ExprVisitor::visit(self.visitor, expr)?; + ExprVisitor::visit(self.visitor, *expr, self.arena)?; } } Ok(()) @@ -387,20 +389,21 @@ pub(crate) mod tests { use crate::planner::operator::join::{JoinOperator, JoinType}; use crate::planner::operator::mark_apply::MarkApplyOperator; use crate::planner::operator::set_membership::SetMembershipKind; + use crate::planner::ExprRef; use crate::planner::{Childrens, LogicalPlan}; use crate::types::index::{IndexInfo, IndexMetaRef, IndexType}; use crate::types::value::DataValue; use crate::types::LogicalType; - fn index_info() -> IndexInfo { + fn index_info(sort_expr: ExprRef, residual_predicate: ExprRef) -> IndexInfo { IndexInfo { meta: IndexMetaRef::new(0), sort_option: SortOption::OrderBy { - fields: vec![SortField::from(ScalarExpression::from(10_i32))], + fields: vec![SortField::from(sort_expr)], ignore_prefix_len: 0, }, lookup: None, - residual_predicate: Some(11_i32.into()), + residual_predicate: Some(residual_predicate), covered_deserializers: None, cover_mapping: None, sort_elimination_hint: None, @@ -408,17 +411,23 @@ pub(crate) mod tests { } } - pub(crate) fn all_operators() -> Result, DatabaseError> { + pub(crate) fn all_operators( + arena: &mut crate::planner::PlanArena, + ) -> Result, DatabaseError> { + let expressions = (0_i32..=18) + .map(|value| arena.alloc_expression(ScalarExpression::from(value))) + .collect::>(); + let expr = |value: usize| expressions[value]; let column_ref = ColumnRef::new(0); let column = ColumnCatalog::new( "value".to_string(), false, - ColumnDesc::new(LogicalType::Integer, None, false, Some(12_i32.into()))?, + ColumnDesc::new(LogicalType::Integer, None, false, Some(expr(12)))?, ); - let mut mark_apply = MarkApplyOperator::new_exists(column_ref, vec![3_i32.into()]); - mark_apply.set_parameterized_probe(Some(4_i32.into())); + let mut mark_apply = MarkApplyOperator::new_exists(column_ref, vec![expr(3)]); + mark_apply.set_parameterized_probe(Some(expr(4))); let table_function = TableFunction { - args: vec![8_i32.into()], + args: vec![expr(8)], catalog: TableFunctionCatalog { schema: Vec::new(), inner: ArcTableFunctionImpl(Numbers::new()), @@ -427,47 +436,47 @@ pub(crate) mod tests { let operators = vec![ Operator::Dummy, Operator::Aggregate(AggregateOperator { - groupby_exprs: vec![1_i32.into()], - agg_calls: vec![2_i32.into()], + groupby_exprs: vec![expr(1)], + agg_calls: vec![expr(2)], is_distinct: false, force_spill: false, }), Operator::ScalarApply(ScalarApplyOperator), Operator::MarkApply(mark_apply), Operator::Filter(FilterOperator { - predicate: 5_i32.into(), + predicate: expr(5), is_optimized: false, having: false, }), Operator::Join(JoinOperator { on: JoinCondition::On { - on: vec![(6_i32.into(), 7_i32.into())], - filter: Some(8_i32.into()), + on: vec![(expr(6), expr(7))], + filter: Some(expr(8)), }, join_type: JoinType::Inner, force_nested_loop: false, }), Operator::Project(ProjectOperator { - exprs: vec![9_i32.into()], + exprs: vec![expr(9)], }), Operator::ScalarSubquery(ScalarSubqueryOperator), Operator::TableScan(TableScanOperator { table_name: "t1".into(), columns: vec![column_ref], limit: (None, None), - index_infos: vec![index_info()], + index_infos: vec![index_info(expr(10), expr(11))], with_pk: false, }), Operator::FunctionScan(FunctionScanOperator { table_function }), Operator::Sort(SortOperator { - sort_fields: vec![SortField::from(ScalarExpression::from(13_i32))], + sort_fields: vec![SortField::from(expr(13))], }), Operator::Limit(LimitOperator { offset: None, limit: Some(1), }), Operator::TopK(TopKOperator { - sort_fields: vec![SortField::from(ScalarExpression::from(14_i32))], + sort_fields: vec![SortField::from(expr(14))], limit: 1, offset: None, }), @@ -476,10 +485,7 @@ pub(crate) mod tests { schema_ref: vec![column_ref], }), Operator::Window(window::WindowOperator { - sort_fields: vec![ - SortField::from(ScalarExpression::from(17_i32)), - SortField::from(ScalarExpression::from(18_i32)), - ], + sort_fields: vec![SortField::from(expr(17)), SortField::from(expr(18))], partition_by_len: 1, functions: vec![WindowFunction { kind: WindowFunctionKind::RowNumber, @@ -510,7 +516,7 @@ pub(crate) mod tests { }), Operator::Update(UpdateOperator { table_name: "t1".into(), - value_exprs: vec![(column_ref, 15_i32.into())], + value_exprs: vec![(column_ref, expr(15))], }), Operator::Delete(DeleteOperator { table_name: "t1".into(), @@ -531,7 +537,7 @@ pub(crate) mod tests { old_column_name: "value".to_string(), new_column_name: "value".to_string(), data_type: LogicalType::Integer, - default_change: DefaultChange::Set(16_i32.into()), + default_change: DefaultChange::Set(expr(16)), not_null_change: NotNullChange::NoChange, }), Operator::DropColumn(DropColumnOperator { @@ -612,10 +618,14 @@ pub(crate) mod tests { #[derive(Default)] struct ExpressionCounter(usize); - impl<'a> ExprVisitor<'a> for ExpressionCounter { - fn visit(&mut self, expr: &'a ScalarExpression) -> Result<(), DatabaseError> { + impl ExprVisitor> for ExpressionCounter { + fn visit( + &mut self, + expr: ExprRef, + arena: &crate::planner::PlanArena<'_>, + ) -> Result<(), DatabaseError> { self.0 += 1; - walk_expr(self, expr) + walk_expr(self, expr, arena) } } @@ -624,13 +634,15 @@ pub(crate) mod tests { struct NoopVisitor; impl OperatorVisitor<'_> for NoopVisitor {} - let operators = all_operators()?; + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = crate::planner::PlanArena::new(&table_arena); + let operators = all_operators(&mut arena)?; for operator in &operators { NoopVisitor.visit_operator(operator)?; } let mut counter = ExpressionCounter::default(); - let mut visitor = OperatorExprVisitor::new(&mut counter); + let mut visitor = OperatorExprVisitor::new(&mut counter, &arena); for operator in &operators { visitor.visit_operator(operator)?; } diff --git a/src/planner/operator/visitor_mut.rs b/src/planner/operator/visitor_mut.rs index 864f36f6..f8ac9c8a 100644 --- a/src/planner/operator/visitor_mut.rs +++ b/src/planner/operator/visitor_mut.rs @@ -16,6 +16,7 @@ use super::alter_table::change_column::DefaultChange; use super::*; use crate::errors::DatabaseError; use crate::expression::visitor_mut::ExprVisitorMut; +use crate::planner::PlanArena; pub trait OperatorVisitorMut<'a>: Sized { fn visit_operator(&mut self, operator: &'a mut Operator) -> Result<(), DatabaseError> { @@ -211,46 +212,47 @@ pub trait OperatorVisitorMut<'a>: Sized { } } -pub struct OperatorExprVisitorMut<'a, V> { +pub struct OperatorExprVisitorMut<'a, 'arena, V> { visitor: &'a mut V, + arena: &'a mut PlanArena<'arena>, } -impl<'a, V> OperatorExprVisitorMut<'a, V> { - pub fn new(visitor: &'a mut V) -> Self { - Self { visitor } +impl<'a, 'arena, V> OperatorExprVisitorMut<'a, 'arena, V> { + pub fn new(visitor: &'a mut V, arena: &'a mut PlanArena<'arena>) -> Self { + Self { visitor, arena } } } -impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMut<'_, V> { +impl<'a, V: ExprVisitorMut> OperatorVisitorMut<'a> for OperatorExprVisitorMut<'_, '_, V> { fn visit_aggregate(&mut self, op: &'a mut AggregateOperator) -> Result<(), DatabaseError> { for expr in op.agg_calls.iter_mut().chain(&mut op.groupby_exprs) { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } fn visit_mark_apply(&mut self, op: &'a mut MarkApplyOperator) -> Result<(), DatabaseError> { for expr in &mut op.predicates { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } if let Some(expr) = &mut op.parameterized_probe { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } fn visit_filter(&mut self, op: &'a mut FilterOperator) -> Result<(), DatabaseError> { - ExprVisitorMut::visit(self.visitor, &mut op.predicate) + ExprVisitorMut::visit(self.visitor, &mut op.predicate, self.arena) } fn visit_join(&mut self, op: &'a mut JoinOperator) -> Result<(), DatabaseError> { if let JoinCondition::On { on, filter } = &mut op.on { for (left_expr, right_expr) in on { - ExprVisitorMut::visit(self.visitor, left_expr)?; - ExprVisitorMut::visit(self.visitor, right_expr)?; + ExprVisitorMut::visit(self.visitor, left_expr, self.arena)?; + ExprVisitorMut::visit(self.visitor, right_expr, self.arena)?; } if let Some(expr) = filter { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } } Ok(()) @@ -258,7 +260,7 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu fn visit_project(&mut self, op: &'a mut ProjectOperator) -> Result<(), DatabaseError> { for expr in &mut op.exprs { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } @@ -267,11 +269,11 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu for index_info in &mut op.index_infos { if let SortOption::OrderBy { fields, .. } = &mut index_info.sort_option { for field in fields { - ExprVisitorMut::visit(self.visitor, &mut field.expr)?; + ExprVisitorMut::visit(self.visitor, &mut field.expr, self.arena)?; } } if let Some(expr) = &mut index_info.residual_predicate { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } } Ok(()) @@ -282,14 +284,14 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu op: &'a mut FunctionScanOperator, ) -> Result<(), DatabaseError> { for expr in &mut op.table_function.args { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } fn visit_sort(&mut self, op: &'a mut SortOperator) -> Result<(), DatabaseError> { for field in &mut op.sort_fields { - ExprVisitorMut::visit(self.visitor, &mut field.expr)?; + ExprVisitorMut::visit(self.visitor, &mut field.expr, self.arena)?; } Ok(()) } @@ -305,28 +307,28 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu .flat_map(|function| &mut function.args), ) { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } fn visit_top_k(&mut self, op: &'a mut TopKOperator) -> Result<(), DatabaseError> { for field in &mut op.sort_fields { - ExprVisitorMut::visit(self.visitor, &mut field.expr)?; + ExprVisitorMut::visit(self.visitor, &mut field.expr, self.arena)?; } Ok(()) } fn visit_update(&mut self, op: &'a mut UpdateOperator) -> Result<(), DatabaseError> { for (_, expr) in &mut op.value_exprs { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } fn visit_add_column(&mut self, op: &'a mut AddColumnOperator) -> Result<(), DatabaseError> { if let Some(expr) = &mut op.column.desc_mut().default { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } @@ -336,7 +338,7 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu op: &'a mut ChangeColumnOperator, ) -> Result<(), DatabaseError> { if let DefaultChange::Set(expr) = &mut op.default_change { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } Ok(()) } @@ -344,7 +346,7 @@ impl<'a, V: ExprVisitorMut<'a>> OperatorVisitorMut<'a> for OperatorExprVisitorMu fn visit_create_table(&mut self, op: &'a mut CreateTableOperator) -> Result<(), DatabaseError> { for column in &mut op.columns { if let Some(expr) = &mut column.desc_mut().default { - ExprVisitorMut::visit(self.visitor, expr)?; + ExprVisitorMut::visit(self.visitor, expr, self.arena)?; } } Ok(()) @@ -408,8 +410,12 @@ mod tests { struct IncrementConstants(usize); - impl ExprVisitorMut<'_> for IncrementConstants { - fn visit_constant(&mut self, value: &mut DataValue) -> Result<(), DatabaseError> { + impl ExprVisitorMut for IncrementConstants { + fn visit_constant( + &mut self, + value: &mut DataValue, + _arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { if let DataValue::Int32(value) = value { *value += 1; } @@ -423,14 +429,16 @@ mod tests { struct NoopVisitor; impl OperatorVisitorMut<'_> for NoopVisitor {} - let mut operators = all_operators()?; + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let mut operators = all_operators(&mut arena)?; for operator in &mut operators { NoopVisitor.visit_operator(operator)?; } let mut counter = IncrementConstants(0); { - let mut visitor = OperatorExprVisitorMut::new(&mut counter); + let mut visitor = OperatorExprVisitorMut::new(&mut counter, &mut arena); for operator in &mut operators { visitor.visit_operator(operator)?; } diff --git a/src/planner/operator/window.rs b/src/planner/operator/window.rs index a5506937..4483130b 100644 --- a/src/planner/operator/window.rs +++ b/src/planner/operator/window.rs @@ -14,11 +14,10 @@ use crate::catalog::ColumnRef; use crate::expression::window::WindowFunction; -use crate::iter_ext::Itertools; use crate::planner::operator::sort::SortField; use crate::planner::operator::SortOption; +use crate::planner::{fmt_explain_list, Explain, PlanArena}; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct WindowOperator { @@ -41,36 +40,36 @@ impl WindowOperator { } } -impl fmt::Display for WindowOperator { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { +impl Explain for WindowOperator { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { let (partition_by, order_by) = self.sort_fields.split_at(self.partition_by_len); - write!( - f, - "Window [{}]", - self.functions - .iter() - .map(|expr| format!("{expr:?}")) - .join(", ") - )?; + f.write_str("Window [")?; + for (index, function) in self.functions.iter().enumerate() { + if index > 0 { + f.write_str(", ")?; + } + write!(f, "WindowFunction {{ kind: {:?}, args: [", function.kind)?; + fmt_explain_list(&function.args, ", ", arena, f)?; + write!(f, "], ty: {:?} }}", function.ty)?; + } + f.write_str("]")?; if !self.sort_fields.is_empty() { - write!(f, " ->")?; + f.write_str(" ->")?; } if !partition_by.is_empty() { - write!( - f, - " Partition By [{}]", - partition_by - .iter() - .map(|field| field.expr.to_string()) - .join(", ") - )?; + f.write_str(" Partition By [")?; + for (index, field) in partition_by.iter().enumerate() { + if index > 0 { + f.write_str(", ")?; + } + field.expr.fmt(arena, f)?; + } + f.write_str("]")?; } if !order_by.is_empty() { - write!( - f, - " Order By [{}]", - order_by.iter().map(ToString::to_string).join(", ") - )?; + f.write_str(" Order By [")?; + fmt_explain_list(order_by, ", ", arena, f)?; + f.write_str("]")?; } Ok(()) } @@ -82,13 +81,13 @@ mod tests { use super::*; use crate::expression::window::WindowFunctionKind; use crate::expression::ScalarExpression; - use crate::planner::TableArena; + use crate::planner::{ExprRef, PlanArena, TableArena, TableArenaCell}; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::rocksdb::RocksTransaction; use crate::types::LogicalType; use std::io::{Cursor, Seek, SeekFrom}; - fn operator(partition_by: Vec, order_by: Vec) -> WindowOperator { + fn operator(partition_by: Vec, order_by: Vec) -> WindowOperator { let partition_by_len = partition_by.len(); WindowOperator { sort_fields: partition_by @@ -107,32 +106,40 @@ mod tests { } #[test] - fn display_window_spec() { + fn explain_window_spec() { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let one = arena.alloc_expression(ScalarExpression::from(1)); + let two = arena.alloc_expression(ScalarExpression::from(2)); let function = "Window [WindowFunction { kind: RowNumber, args: [], ty: Bigint }]"; - assert_eq!(operator(Vec::new(), Vec::new()).to_string(), function); assert_eq!( - operator(vec![1.into()], Vec::new()).to_string(), - format!("{function} -> Partition By [1]") + operator(Vec::new(), Vec::new()).explain(&arena).to_string(), + function ); assert_eq!( - operator(Vec::new(), vec![ScalarExpression::from(2).desc()]).to_string(), - format!("{function} -> Order By [2 Desc Nulls Last]") + operator(vec![one], Vec::new()).explain(&arena).to_string(), + function.to_owned() + " -> Partition By [1]" ); assert_eq!( - operator(vec![1.into()], vec![ScalarExpression::from(2).desc()]).to_string(), - format!("{function} -> Partition By [1] Order By [2 Desc Nulls Last]") + operator(Vec::new(), vec![SortField::from(two).desc()]) + .explain(&arena) + .to_string(), + function.to_owned() + " -> Order By [2 Desc Nulls Last]" + ); + assert_eq!( + operator(vec![one], vec![SortField::from(two).desc()]) + .explain(&arena) + .to_string(), + function.to_owned() + " -> Partition By [1] Order By [2 Desc Nulls Last]" ); assert_eq!( operator(Vec::new(), Vec::new()).sort_option(), SortOption::Follow ); assert_eq!( - operator(vec![1.into()], vec![ScalarExpression::from(2).desc()]).sort_option(), + operator(vec![one], vec![SortField::from(two).desc()]).sort_option(), SortOption::OrderBy { - fields: vec![ - ScalarExpression::from(1).asc(), - ScalarExpression::from(2).desc(), - ], + fields: vec![SortField::from(one).asc(), SortField::from(two).desc()], ignore_prefix_len: 0, } ); @@ -140,10 +147,13 @@ mod tests { #[test] fn serialization_roundtrip() -> Result<(), crate::errors::DatabaseError> { - let source = operator(vec![1.into()], vec![ScalarExpression::from(2).desc()]); + let mut arena = TableArena::default(); + let source = operator( + vec![arena.alloc_expression(ScalarExpression::from(1))], + vec![SortField::from(arena.alloc_expression(ScalarExpression::from(2))).desc()], + ); let mut cursor = Cursor::new(Vec::new()); let mut reference_tables = ReferenceTables::new(); - let arena = TableArena::default(); source.encode(&mut cursor, false, &mut reference_tables, &arena)?; cursor.seek(SeekFrom::Start(0))?; diff --git a/src/serdes/column.rs b/src/serdes/column.rs index 682c6810..b5a7d2da 100644 --- a/src/serdes/column.rs +++ b/src/serdes/column.rs @@ -209,15 +209,12 @@ pub(crate) mod test { cursor.seek(SeekFrom::Start(0))?; } { + let default = + plan_arena.alloc_expression(ScalarExpression::Constant(DataValue::UInt64(42))); let not_ref_column = plan_arena.alloc_column(ColumnCatalog::new( "c3".to_string(), false, - ColumnDesc::new( - LogicalType::Integer, - None, - false, - Some(ScalarExpression::Constant(DataValue::UInt64(42))), - )?, + ColumnDesc::new(LogicalType::Integer, None, false, Some(default))?, )); not_ref_column.encode(&mut cursor, false, &mut reference_tables, &plan_arena)?; cursor.seek(SeekFrom::Start(0))?; @@ -228,10 +225,16 @@ pub(crate) mod test { &reference_tables, &mut plan_arena, )?; - assert_eq!( - plan_arena.column(decoded), - plan_arena.column(not_ref_column) - ); + let decoded = plan_arena.column(decoded); + let expected = plan_arena.column(not_ref_column); + assert_eq!(decoded.summary(), expected.summary()); + assert_eq!(decoded.nullable(), expected.nullable()); + assert_eq!(decoded.datatype(), expected.datatype()); + assert_eq!(decoded.desc().primary(), expected.desc().primary()); + assert_eq!(decoded.desc().is_unique(), expected.desc().is_unique()); + let decoded_default = decoded.desc().default.unwrap(); + let expected_default = expected.desc().default.unwrap(); + assert!(decoded_default.eq_ignore_colref_pos(expected_default, &plan_arena)); } Ok(()) @@ -311,7 +314,7 @@ pub(crate) mod test { LogicalType::Integer, None, false, - Some(ScalarExpression::Constant(DataValue::UInt64(42))), + Some(arena.alloc_expression(ScalarExpression::Constant(DataValue::UInt64(42)))), )?; desc.encode(&mut cursor, false, &mut reference_tables, &arena)?; cursor.seek(SeekFrom::Start(0))?; @@ -322,7 +325,13 @@ pub(crate) mod test { &reference_tables, &mut arena, )?; - assert_eq!(desc, decode_desc); + assert_eq!(desc.column_datatype, decode_desc.column_datatype); + assert_eq!(desc.primary(), decode_desc.primary()); + assert_eq!(desc.is_unique(), decode_desc.is_unique()); + assert_eq!( + arena.expression(desc.default.unwrap()), + arena.expression(decode_desc.default.unwrap()) + ); Ok(()) } diff --git a/src/serdes/expression.rs b/src/serdes/expression.rs new file mode 100644 index 00000000..c516158c --- /dev/null +++ b/src/serdes/expression.rs @@ -0,0 +1,44 @@ +// Copyright 2024 KipData/KiteSQL +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use crate::errors::DatabaseError; +use crate::expression::ScalarExpression; +use crate::planner::{ExprRef, MetaArena}; +use crate::serdes::{ReferenceDecodeContext, ReferenceSerialization, ReferenceTables}; +use crate::storage::Transaction; +use std::io::{Read, Write}; + +impl ReferenceSerialization for ExprRef { + fn encode( + &self, + writer: &mut W, + is_direct: bool, + reference_tables: &mut ReferenceTables, + arena: &A, + ) -> Result<(), DatabaseError> { + arena + .expression(*self) + .encode(writer, is_direct, reference_tables, arena) + } + + fn decode( + reader: &mut R, + context: Option<&ReferenceDecodeContext<'_, T>>, + reference_tables: &ReferenceTables, + arena: &mut A, + ) -> Result { + let expression = ScalarExpression::decode(reader, context, reference_tables, arena)?; + Ok(arena.alloc_expression(expression)) + } +} diff --git a/src/serdes/mod.rs b/src/serdes/mod.rs index 8b1c0d53..eeddd439 100644 --- a/src/serdes/mod.rs +++ b/src/serdes/mod.rs @@ -20,6 +20,7 @@ mod char_length_units; mod column; mod data_value; mod evaluator; +mod expression; mod function; mod hasher; mod index; diff --git a/src/storage/mod.rs b/src/storage/mod.rs index 9c5d48e8..4fb48959 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -24,7 +24,7 @@ use crate::catalog::{ColumnCatalog, ColumnRef, TableCatalog, TableMeta, TableNam use crate::db::{ScalaFunctions, TableFunctions}; use crate::errors::DatabaseError; use crate::expression::range_detacher::Range; -use crate::expression::ScalarExpression; +use crate::expression::TypeCast; use crate::iter_ext::Itertools; use crate::optimizer::core::cm_sketch::{ CountMinSketch, CountMinSketchPage, COUNT_MIN_SKETCH_STORAGE_PAGE_LEN, @@ -473,24 +473,22 @@ pub trait Transaction: Sized { return Err(DatabaseError::DuplicateColumn(new_column_name.to_string())); } - for column in table.columns().map(|column| plan_arena.column(*column)) { - let mut new_column = ColumnCatalog::clone(column); - if column.name() == old_column_name { + for column_ref in table.columns() { + let mut new_column = plan_arena.column(*column_ref).clone(); + if new_column.name() == old_column_name { found = true; new_column.set_name(new_column_name.to_string()); new_column.desc_mut().column_datatype = new_data_type.clone(); match default_change { DefaultChange::NoChange => { - if let Some(default_expr) = new_column.desc().default.clone() { - new_column.desc_mut().default = Some(ScalarExpression::type_cast( - default_expr, - Cow::Borrowed(new_data_type), - plan_arena, - )?); + if let Some(default_expr) = new_column.desc().default { + new_column.desc_mut().default = Some( + default_expr.type_cast(Cow::Borrowed(new_data_type), plan_arena)?, + ); } } DefaultChange::Set(default_expr) => { - new_column.desc_mut().default = Some(default_expr.clone()); + new_column.desc_mut().default = Some(*default_expr); } DefaultChange::Drop => { new_column.desc_mut().default = None; @@ -558,7 +556,7 @@ pub trait Transaction: Sized { let mut table = self .load_table(table_codec, plan_arena, table_name.clone())? .ok_or(DatabaseError::TableNotFound)?; - if !column.nullable() && column.default_value()?.is_none() { + if !column.nullable() && column.default_value(plan_arena)?.is_none() { return Err(DatabaseError::NeedNullAbleOrDefault); } @@ -2890,7 +2888,7 @@ mod test { let mut transaction = table_state.storage.transaction()?; let mut view_cache = table_state.view_cache.clone(); let mut table_codec = TableCodec::default(); - let view = transaction.create_view(&mut table_codec, &plan_arena, view.clone(), true)?; + let view = transaction.create_view(&mut table_codec, &plan_arena, view, true)?; view_cache.insert(view.name.clone(), view.clone()); assert_eq!( diff --git a/src/storage/table_codec.rs b/src/storage/table_codec.rs index 6921506f..1c0145e5 100644 --- a/src/storage/table_codec.rs +++ b/src/storage/table_codec.rs @@ -1145,7 +1145,7 @@ mod tests { root.histogram_meta().buckets_len() ); - let (sketch_meta, mut sketch_pages) = sketch.clone().into_storage_parts(1); + let (sketch_meta, mut sketch_pages) = sketch.into_storage_parts(1); let sketch_meta_bytes = table_codec.with_statistics_sketch_meta( "t1", 0, @@ -1323,6 +1323,9 @@ mod tests { } else if bytes[i] == b'#' { normalized.push_str("#_"); i += 1; + if bytes[i..].starts_with(b"expr") { + i += b"expr".len(); + } while i < bytes.len() && bytes[i].is_ascii_digit() { i += 1; } diff --git a/src/types/evaluator/cast.rs b/src/types/evaluator/cast.rs index 6c9ed1be..12e7bc74 100644 --- a/src/types/evaluator/cast.rs +++ b/src/types/evaluator/cast.rs @@ -42,7 +42,6 @@ use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; use crate::types::LogicalType; use paste::paste; -use std::borrow::Cow; pub(crate) fn cast_fail(from: LogicalType, to: LogicalType) -> DatabaseError { DatabaseError::CastFail { @@ -530,11 +529,9 @@ macro_rules! build_integer_cast { } pub fn cast_create( - from: Cow<'_, LogicalType>, - to: Cow<'_, LogicalType>, + from: &LogicalType, + to: &LogicalType, ) -> Result { - let from = from.as_ref(); - let to = to.as_ref(); if from == to { return Ok(CastEvaluatorRef::new( cast_pos(from, to), @@ -803,7 +800,7 @@ pub fn cast_create( let evaluators = from_types .iter() .zip(to_types.iter()) - .map(|(from, to)| cast_create(Cow::Borrowed(from), Cow::Borrowed(to))) + .map(|(from, to)| cast_create(from, to)) .collect::, _>>()?; Ok(CastEvaluatorRef::new( cast_pos(from, to), @@ -1160,11 +1157,10 @@ mod test { use ordered_float::OrderedFloat; #[cfg(feature = "decimal")] use rust_decimal::Decimal; - use std::borrow::Cow; use std::io::{Cursor, Seek, SeekFrom}; fn create(from: LogicalType, to: LogicalType) -> Result { - cast_create(Cow::Owned(from), Cow::Owned(to)) + cast_create(&from, &to) } fn utf8(value: &str) -> DataValue { diff --git a/src/types/evaluator/tuple.rs b/src/types/evaluator/tuple.rs index 753c5754..5a54220f 100644 --- a/src/types/evaluator/tuple.rs +++ b/src/types/evaluator/tuple.rs @@ -122,7 +122,6 @@ mod test { use crate::types::evaluator::cast_create; use crate::types::CharLengthUnits; use crate::types::LogicalType; - use std::borrow::Cow; fn tuple(values: Vec) -> DataValue { DataValue::Tuple(values, false) @@ -194,14 +193,11 @@ mod test { #[test] fn test_tuple_cast_eval() { let evaluator = cast_create( - Cow::Owned(LogicalType::Tuple(vec![ + &LogicalType::Tuple(vec![ LogicalType::Integer, LogicalType::Varchar(None, CharLengthUnits::Characters), - ])), - Cow::Owned(LogicalType::Tuple(vec![ - LogicalType::Bigint, - LogicalType::Integer, - ])), + ]), + &LogicalType::Tuple(vec![LogicalType::Bigint, LogicalType::Integer]), ) .unwrap(); diff --git a/src/types/index.rs b/src/types/index.rs index 27eb1432..6e83fee1 100644 --- a/src/types/index.rs +++ b/src/types/index.rs @@ -17,7 +17,8 @@ use crate::errors::DatabaseError; use crate::expression::range_detacher::Range; use crate::expression::ScalarExpression; use crate::planner::operator::SortOption; -use crate::planner::PlanArena; +use crate::planner::Explain; +use crate::planner::{ExprRef, PlanArena}; use crate::types::serialize::TupleValueSerializableImpl; use crate::types::value::DataValue; use crate::types::{ColumnId, LogicalType}; @@ -73,7 +74,7 @@ pub struct IndexInfo { pub(crate) meta: IndexMetaRef, pub(crate) sort_option: SortOption, pub(crate) lookup: Option, - pub(crate) residual_predicate: Option, + pub(crate) residual_predicate: Option, pub(crate) covered_deserializers: Option>, pub(crate) cover_mapping: Option>, pub(crate) sort_elimination_hint: Option, @@ -147,23 +148,17 @@ impl<'a> Index<'a> { } } -impl fmt::Display for IndexInfo { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - write!(f, "{}", self.meta)?; - write!(f, " => ")?; - - if let Some(lookup) = &self.lookup { - match lookup { - IndexLookup::Static(range) => write!(f, "{range}")?, - IndexLookup::Probe => write!(f, "Probe ?")?, - } - } else { - write!(f, "EMPTY")?; +impl Explain for IndexInfo { + fn fmt(&self, arena: &PlanArena<'_>, f: &mut Formatter<'_>) -> fmt::Result { + write!(f, "{} => ", self.meta.explain(arena))?; + match &self.lookup { + Some(IndexLookup::Static(range)) => write!(f, "{range}")?, + Some(IndexLookup::Probe) => f.write_str("Probe ?")?, + None => f.write_str("EMPTY")?, } if self.covered_deserializers.is_some() { - write!(f, " Covered")?; + f.write_str(" Covered")?; } - Ok(()) } } @@ -234,7 +229,7 @@ mod tests { } #[test] - fn test_index_helpers_and_display() { + fn test_index_helpers_and_explain() { let meta_ref = IndexMetaRef::new(3); assert_eq!(meta_ref.pos(), 3); assert_eq!(meta_ref.to_string(), "#3"); @@ -249,14 +244,22 @@ mod tests { let meta = index_meta(); assert_eq!(meta.to_string(), "idx_t"); - assert_eq!(index_info(None).to_string(), "#7 => EMPTY"); - assert_eq!( - index_info(Some(IndexLookup::Probe)).to_string(), - "#7 => Probe ?" - ); + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let index_ref = arena.alloc_index(index_meta()); + + let mut empty = index_info(None); + empty.meta = index_ref; + assert_eq!(empty.explain(&arena).to_string(), "idx_t => EMPTY"); + + let mut probe = index_info(Some(IndexLookup::Probe)); + probe.meta = index_ref; + assert_eq!(probe.explain(&arena).to_string(), "idx_t => Probe ?"); + let mut info = index_info(Some(IndexLookup::Static(Range::Eq(DataValue::Int32(1))))); + info.meta = index_ref; info.covered_deserializers = Some(vec![LogicalType::Integer.serializable()]); - assert_eq!(info.to_string(), "#7 => 1 Covered"); + assert_eq!(info.explain(&arena).to_string(), "idx_t => 1 Covered"); let value = DataValue::Int32(1); let index = Index::new(9, &value, IndexType::Unique); diff --git a/src/types/value.rs b/src/types/value.rs index e10a8509..44922b1c 100644 --- a/src/types/value.rs +++ b/src/types/value.rs @@ -26,7 +26,6 @@ use chrono::{ use ordered_float::OrderedFloat; #[cfg(feature = "decimal")] use rust_decimal::Decimal; -use std::borrow::Cow; use std::cmp::Ordering; use std::fmt::Formatter; use std::hash::Hash; @@ -1310,7 +1309,7 @@ impl DataValue { to_varchar(value, *len, *unit) } (value, _) => { - let evaluator = cast_create(Cow::Owned(from), Cow::Borrowed(to))?; + let evaluator = cast_create(&from, to)?; evaluator.eval(&value) } } diff --git a/tests/macros-test/src/main.rs b/tests/macros-test/src/main.rs index 37b84c66..5d9d1746 100644 --- a/tests/macros-test/src/main.rs +++ b/tests/macros-test/src/main.rs @@ -2534,9 +2534,9 @@ mod test { }) .collect::>() .join("\n"); - assert!( - explain_plan.contains("IndexScan By #") && explain_plan.contains("Covered"), - "unexpected explain plan: {explain_plan}" + assert_eq!( + explain_plan, + "Projection [users.age] [Project => (Sort Option: Follow)] TableScan users -> [users.age] [IndexScan By users_age_index => 1050 Covered => (Sort Option: OrderBy: (users.age Asc Nulls Last) ignore_prefix_len: 0)]" ); database @@ -2946,10 +2946,10 @@ mod test { .project_scalar(User::name())? .finish() })?; - assert!(plan.contains("Projection")); - assert!(plan.contains("Filter (")); - assert!(plan.contains(" = 1")); - assert!(plan.contains("TableScan users -> [#")); + assert_eq!( + plan, + "Projection [users.user_name] [Project => (Sort Option: Follow)] TableScan users -> [users.id, users.user_name] [IndexScan By pk_index => 1 => (Sort Option: OrderBy: (users.id Asc Nulls Last) ignore_prefix_len: 0)]" + ); let set_plan = database.explain(|ctx| { ctx.union( @@ -2958,10 +2958,10 @@ mod test { |ctx| ctx.from::()?.project_scalar(Wallet::id())?.finish(), ) })?; - assert!(set_plan.contains("Aggregate")); - assert!(set_plan.contains("Union: [#")); - assert!(set_plan.contains("TableScan users -> [#")); - assert!(set_plan.contains("TableScan wallets -> [#")); + assert_eq!( + set_plan, + "Aggregate [] -> Group By [users.id] [HashAggregate => (Sort Option: None)] Union: [users.id] Projection [users.id] [Project => (Sort Option: Follow)] TableScan users -> [users.id] [SeqScan => (Sort Option: None)] Projection [wallets.id] [Project => (Sort Option: Follow)] TableScan wallets -> [wallets.id] [SeqScan => (Sort Option: None)]" + ); let mut tx = database.new_transaction()?; let tx_tables = tx.show_tables()?.collect::, _>>()?; @@ -3047,15 +3047,18 @@ mod test { #[test] fn test_scala_function() -> Result<(), DatabaseError> { let function = MyScalaFunction::new(); + let table_arena = kite_sql::planner::TableArenaCell::default(); + let mut arena = kite_sql::planner::PlanArena::new(&table_arena); let sum = function.eval( &[ - ScalarExpression::Constant(DataValue::Int8(1)), - ScalarExpression::Constant(DataValue::Utf8 { + arena.alloc_expression(ScalarExpression::Constant(DataValue::Int8(1))), + arena.alloc_expression(ScalarExpression::Constant(DataValue::Utf8 { value: "1".to_string(), ty: Utf8Type::Variable(None), unit: CharLengthUnits::Characters, - }), + })), ], + &arena, None, )?; @@ -3075,7 +3078,12 @@ mod test { #[test] fn test_table_function() -> Result<(), DatabaseError> { let function = MyTableFunction::new(); - let mut numbers = function.eval(&[ScalarExpression::Constant(DataValue::Int8(2))])?; + let table_arena = kite_sql::planner::TableArenaCell::default(); + let mut arena = kite_sql::planner::PlanArena::new(&table_arena); + let mut numbers = function.eval( + &[arena.alloc_expression(ScalarExpression::Constant(DataValue::Int8(2)))], + &arena, + )?; println!("{:?}", function); diff --git a/tests/slt/cte.slt b/tests/slt/cte.slt index f3a245d5..f35c9ac4 100644 --- a/tests/slt/cte.slt +++ b/tests/slt/cte.slt @@ -99,7 +99,7 @@ EXPLAIN WITH RECURSIVE recursive_cte(value) AS ( ) SELECT value FROM recursive_cte ---- -Projection [#4] [Project => (Sort Option: Follow)] Recursive CTE: [#4] Projection [(#3) as (#4)] [Project => (Sort Option: Follow)] Projection [1] [Project => (Sort Option: Follow)] Dummy [Dummy => (Sort Option: None)] Projection [(#4 + 1)] [Project => (Sort Option: Follow)] Filter (#4 < 3), Is Having: false [Filter => (Sort Option: Follow)] Recursive Scan: [#4] +Projection [recursive_cte.value] [Project => (Sort Option: Follow)] Recursive CTE: [recursive_cte.value] Projection [(1) as (recursive_cte.value)] [Project => (Sort Option: Follow)] Projection [1] [Project => (Sort Option: Follow)] Dummy [Dummy => (Sort Option: None)] Projection [(recursive_cte.value + 1)] [Project => (Sort Option: Follow)] Filter (recursive_cte.value < 3), Is Having: false [Filter => (Sort Option: Follow)] Recursive Scan: [recursive_cte.value] query I WITH RECURSIVE recursive_cte(value) AS ( diff --git a/tests/slt/join.slt b/tests/slt/join.slt index 439ab503..b4aafd61 100644 --- a/tests/slt/join.slt +++ b/tests/slt/join.slt @@ -31,7 +31,7 @@ select a, b, c, d from x join y on a = c; query T explain select /*+ FORCE_NEST_LOOP_JOIN */ a, b, c, d from x join y on a = c; ---- -Projection [#10, #11, #12, #13] [Project => (Sort Option: Follow)] Inner Join On #2 = #5 [NestLoopJoin => (Sort Option: None)] TableScan x -> [#2, #3] [SeqScan => (Sort Option: None)] TableScan y -> [#5, #6] [SeqScan => (Sort Option: None)] +Projection [x.a, x.b, y.c, y.d] [Project => (Sort Option: Follow)] Inner Join On x.a = y.c [NestLoopJoin => (Sort Option: None)] TableScan x -> [x.a, x.b] [SeqScan => (Sort Option: None)] TableScan y -> [y.c, y.d] [SeqScan => (Sort Option: None)] query IIII select a, b, c, d, e, f from x join y on a = c and c < 5 join z on e = a and f = 5; diff --git a/tests/slt/parameterized_subquery.slt b/tests/slt/parameterized_subquery.slt new file mode 100644 index 00000000..58d4f47e --- /dev/null +++ b/tests/slt/parameterized_subquery.slt @@ -0,0 +1,173 @@ +# Correlated IN/NOT IN should use a parameterized index probe. +statement ok +create table in_outer(id int primary key, a int); + +statement ok +create table in_inner(id int primary key, v int); + +statement ok +create table in_inner_nn(id int primary key, v int); + +statement ok +create index in_inner_v_index on in_inner(v); + +statement ok +create index in_inner_nn_v_index on in_inner_nn(v); + +statement ok +insert into in_outer values (0, null), (1, 1), (2, 2), (3, 3); + +statement ok +insert into in_inner values (0, 2), (1, null); + +statement ok +insert into in_inner_nn values (0, 2); + +query T +explain select id from in_outer where a in (select v from in_inner where in_inner.v = in_outer.a); +---- +Projection [in_outer.id] [Project => (Sort Option: Follow)] Filter _temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer -> [in_outer.id, in_outer.a] [SeqScan => (Sort Option: None)] Projection [in_inner.v] [Project => (Sort Option: Follow)] TableScan in_inner -> [in_inner.v] [IndexScan By in_inner_v_index => Probe ? => (Sort Option: OrderBy: (in_inner.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer where a not in (select v from in_inner where in_inner.v = in_outer.a); +---- +Projection [in_outer.id] [Project => (Sort Option: Follow)] Filter !_temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer -> [in_outer.id, in_outer.a] [SeqScan => (Sort Option: None)] Projection [in_inner.v] [Project => (Sort Option: Follow)] TableScan in_inner -> [in_inner.v] [IndexScan By in_inner_v_index => Probe ? => (Sort Option: OrderBy: (in_inner.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer where a in (select v from in_inner_nn where in_inner_nn.v = in_outer.a); +---- +Projection [in_outer.id] [Project => (Sort Option: Follow)] Filter _temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer -> [in_outer.id, in_outer.a] [SeqScan => (Sort Option: None)] Projection [in_inner_nn.v] [Project => (Sort Option: Follow)] TableScan in_inner_nn -> [in_inner_nn.v] [IndexScan By in_inner_nn_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_nn.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer where a not in (select v from in_inner_nn where in_inner_nn.v = in_outer.a); +---- +Projection [in_outer.id] [Project => (Sort Option: Follow)] Filter !_temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer -> [in_outer.id, in_outer.a] [SeqScan => (Sort Option: None)] Projection [in_inner_nn.v] [Project => (Sort Option: Follow)] TableScan in_inner_nn -> [in_inner_nn.v] [IndexScan By in_inner_nn_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_nn.v Asc Nulls Last) ignore_prefix_len: 0)] + +query I rowsort +select id from in_outer where a in (select v from in_inner where in_inner.v = in_outer.a) order by id; +---- +2 + +query I rowsort +select id from in_outer where a not in (select v from in_inner where in_inner.v = in_outer.a) order by id; +---- +0 +1 +3 + +query I rowsort +select id from in_outer where a in (select v from in_inner_nn where in_inner_nn.v = in_outer.a) order by id; +---- +2 + +query I rowsort +select id from in_outer where a not in (select v from in_inner_nn where in_inner_nn.v = in_outer.a) order by id; +---- +0 +1 +3 + +# Correlated IN/NOT IN with an additional correlated predicate should still +# retain the parameterized probe and the projected predicate column. +statement ok +create table in_outer_flag(id int primary key, a int, b int); + +statement ok +create table in_inner_flag(id int primary key, v int, flag int); + +statement ok +create table in_inner_flag_nn(id int primary key, v int, flag int); + +statement ok +create index in_inner_flag_v_index on in_inner_flag(v); + +statement ok +create index in_inner_flag_nn_v_index on in_inner_flag_nn(v); + +statement ok +insert into in_outer_flag values (0, null, 1), (1, 1, 1), (2, 2, 1), (3, 3, 1); + +statement ok +insert into in_inner_flag values (0, 2, 1), (1, null, 1); + +statement ok +insert into in_inner_flag_nn values (0, 2, 1); + +query T +explain select id from in_outer_flag where a in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b); +---- +Projection [in_outer_flag.id] [Project => (Sort Option: Follow)] Filter _temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer_flag -> [in_outer_flag.id, in_outer_flag.a, in_outer_flag.b] [SeqScan => (Sort Option: None)] Projection [in_inner_flag.v, in_inner_flag.flag] [Project => (Sort Option: Follow)] TableScan in_inner_flag -> [in_inner_flag.v, in_inner_flag.flag] [IndexScan By in_inner_flag_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_flag.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer_flag where a not in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b); +---- +Projection [in_outer_flag.id] [Project => (Sort Option: Follow)] Filter !_temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer_flag -> [in_outer_flag.id, in_outer_flag.a, in_outer_flag.b] [SeqScan => (Sort Option: None)] Projection [in_inner_flag.v, in_inner_flag.flag] [Project => (Sort Option: Follow)] TableScan in_inner_flag -> [in_inner_flag.v, in_inner_flag.flag] [IndexScan By in_inner_flag_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_flag.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer_flag where a in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b); +---- +Projection [in_outer_flag.id] [Project => (Sort Option: Follow)] Filter _temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer_flag -> [in_outer_flag.id, in_outer_flag.a, in_outer_flag.b] [SeqScan => (Sort Option: None)] Projection [in_inner_flag_nn.v, in_inner_flag_nn.flag] [Project => (Sort Option: Follow)] TableScan in_inner_flag_nn -> [in_inner_flag_nn.v, in_inner_flag_nn.flag] [IndexScan By in_inner_flag_nn_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_flag_nn.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from in_outer_flag where a not in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b); +---- +Projection [in_outer_flag.id] [Project => (Sort Option: Follow)] Filter !_temp_table_1_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkAnyApply TableScan in_outer_flag -> [in_outer_flag.id, in_outer_flag.a, in_outer_flag.b] [SeqScan => (Sort Option: None)] Projection [in_inner_flag_nn.v, in_inner_flag_nn.flag] [Project => (Sort Option: Follow)] TableScan in_inner_flag_nn -> [in_inner_flag_nn.v, in_inner_flag_nn.flag] [IndexScan By in_inner_flag_nn_v_index => Probe ? => (Sort Option: OrderBy: (in_inner_flag_nn.v Asc Nulls Last) ignore_prefix_len: 0)] + +query I rowsort +select id from in_outer_flag where a in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b) order by id; +---- +2 + +query I rowsort +select id from in_outer_flag where a not in (select v from in_inner_flag where in_inner_flag.flag = in_outer_flag.b) order by id; +---- + +query I rowsort +select id from in_outer_flag where a in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b) order by id; +---- +2 + +query I rowsort +select id from in_outer_flag where a not in (select v from in_inner_flag_nn where in_inner_flag_nn.flag = in_outer_flag.b) order by id; +---- +1 +3 + +# Correlated EXISTS/NOT EXISTS should use a parameterized index probe while +# retaining non-index correlated predicates. +statement ok +create table exists_outer(id int primary key, a int, b int); + +statement ok +create table exists_inner(id int primary key, v int, flag int); + +statement ok +create index exists_inner_v_index on exists_inner(v); + +statement ok +insert into exists_outer values (0, 1, 1), (1, 1, 2), (2, 2, null), (3, 3, 1); + +statement ok +insert into exists_inner values (0, 1, 1), (1, 1, null), (2, 2, 1); + +query T +explain select id from exists_outer where exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b); +---- +Projection [exists_outer.id] [Project => (Sort Option: Follow)] Filter _temp_table_0_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkExistsApply TableScan exists_outer -> [exists_outer.id, exists_outer.a, exists_outer.b] [SeqScan => (Sort Option: None)] TableScan exists_inner -> [exists_inner.id, exists_inner.v, exists_inner.flag] [IndexScan By exists_inner_v_index => Probe ? => (Sort Option: OrderBy: (exists_inner.v Asc Nulls Last) ignore_prefix_len: 0)] + +query T +explain select id from exists_outer where not exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b); +---- +Projection [exists_outer.id] [Project => (Sort Option: Follow)] Filter !_temp_table_0_.true, Is Having: false [Filter => (Sort Option: Follow)] MarkExistsApply TableScan exists_outer -> [exists_outer.id, exists_outer.a, exists_outer.b] [SeqScan => (Sort Option: None)] TableScan exists_inner -> [exists_inner.id, exists_inner.v, exists_inner.flag] [IndexScan By exists_inner_v_index => Probe ? => (Sort Option: OrderBy: (exists_inner.v Asc Nulls Last) ignore_prefix_len: 0)] + +query I rowsort +select id from exists_outer where exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b) order by id; +---- +0 + +query I rowsort +select id from exists_outer where not exists (select 1 from exists_inner where exists_inner.v = exists_outer.a and exists_inner.flag = exists_outer.b) order by id; +---- +1 +2 +3 diff --git a/tests/slt/stream_distinct_explain.slt b/tests/slt/stream_distinct_explain.slt index 5bbec0a9..daa965c7 100644 --- a/tests/slt/stream_distinct_explain.slt +++ b/tests/slt/stream_distinct_explain.slt @@ -14,19 +14,19 @@ analyze table distinct_t; query T explain select distinct c1 from distinct_t where c1 < 10 and c1 > 0; ---- -Projection [#2] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [#2] [StreamDistinct => (Sort Option: Follow)] TableScan distinct_t -> [#2] [IndexScan By #1 => (0, 10) Covered => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [distinct_t.c1] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [distinct_t.c1] [StreamDistinct => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1] [IndexScan By distinct_t_c1_index => (0, 10) Covered => (Sort Option: OrderBy: (distinct_t.c1 Asc Nulls Last) ignore_prefix_len: 0)] # stream aggregate query T explain select c1, count(c2) from distinct_t where c1 < 10 and c1 > 0 group by c1; ---- -Projection [#2, #4] [Project => (Sort Option: Follow)] Aggregate [Count(#3)] -> Group By [#2] [StreamAggregate => (Sort Option: Follow)] TableScan distinct_t -> [#2, #3] [IndexScan By #1 => (0, 10) => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [distinct_t.c1, Count(distinct_t.c2)] [Project => (Sort Option: Follow)] Aggregate [Count(distinct_t.c2)] -> Group By [distinct_t.c1] [StreamAggregate => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1, distinct_t.c2] [IndexScan By distinct_t_c1_index => (0, 10) => (Sort Option: OrderBy: (distinct_t.c1 Asc Nulls Last) ignore_prefix_len: 0)] # forced spill reuses ordered input query T explain select /*+ FORCE_AGG_SPILL */ c1, count(c2) from distinct_t where c1 < 10 and c1 > 0 group by c1; ---- -Projection [#2, #4] [Project => (Sort Option: Follow)] Aggregate [Count(#3)] -> Group By [#2] [StreamAggregate => (Sort Option: Follow)] TableScan distinct_t -> [#2, #3] [IndexScan By #1 => (0, 10) => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [distinct_t.c1, Count(distinct_t.c2)] [Project => (Sort Option: Follow)] Aggregate [Count(distinct_t.c2)] -> Group By [distinct_t.c1] [StreamAggregate => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1, distinct_t.c2] [IndexScan By distinct_t_c1_index => (0, 10) => (Sort Option: OrderBy: (distinct_t.c1 Asc Nulls Last) ignore_prefix_len: 0)] statement ok drop index distinct_t.distinct_t_c1_index; @@ -35,19 +35,19 @@ drop index distinct_t.distinct_t_c1_index; query T explain select distinct c1 from distinct_t where c1 < 10 and c1 > 0; ---- -Projection [#2] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [#2] [HashAggregate => (Sort Option: None)] Filter ((#2 < 10) && (#2 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [#2] [SeqScan => (Sort Option: None)] +Projection [distinct_t.c1] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [distinct_t.c1] [HashAggregate => (Sort Option: None)] Filter ((distinct_t.c1 < 10) && (distinct_t.c1 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1] [SeqScan => (Sort Option: None)] # forced spill aggregate query T explain select /*+ FORCE_AGG_SPILL */ c1, count(c2) from distinct_t where c1 < 10 and c1 > 0 group by c1; ---- -Projection [#2, #4] [Project => (Sort Option: Follow)] Aggregate [Count(#3)] -> Group By [#2] [StreamAggregate => (Sort Option: Follow)] Sort By #2 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] Filter ((#2 < 10) && (#2 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [#2, #3] [SeqScan => (Sort Option: None)] +Projection [distinct_t.c1, Count(distinct_t.c2)] [Project => (Sort Option: Follow)] Aggregate [Count(distinct_t.c2)] -> Group By [distinct_t.c1] [StreamAggregate => (Sort Option: Follow)] Sort By distinct_t.c1 Asc Nulls Last [Sort => (Sort Option: OrderBy: (distinct_t.c1 Asc Nulls Last) ignore_prefix_len: 0)] Filter ((distinct_t.c1 < 10) && (distinct_t.c1 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1, distinct_t.c2] [SeqScan => (Sort Option: None)] # forced spill sort also satisfies order by query T explain select /*+ FORCE_AGG_SPILL */ c1, count(c2) from distinct_t where c1 < 10 and c1 > 0 group by c1 order by c1; ---- -Projection [#2, #4] [Project => (Sort Option: Follow)] Aggregate [Count(#3)] -> Group By [#2] [StreamAggregate => (Sort Option: Follow)] Sort By #2 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] Filter ((#2 < 10) && (#2 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [#2, #3] [SeqScan => (Sort Option: None)] +Projection [distinct_t.c1, Count(distinct_t.c2)] [Project => (Sort Option: Follow)] Aggregate [Count(distinct_t.c2)] -> Group By [distinct_t.c1] [StreamAggregate => (Sort Option: Follow)] Sort By distinct_t.c1 Asc Nulls Last [Sort => (Sort Option: OrderBy: (distinct_t.c1 Asc Nulls Last) ignore_prefix_len: 0)] Filter ((distinct_t.c1 < 10) && (distinct_t.c1 > 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan distinct_t -> [distinct_t.c1, distinct_t.c2] [SeqScan => (Sort Option: None)] statement ok drop table distinct_t; diff --git a/tests/slt/subquery.slt b/tests/slt/subquery.slt index e6e55741..ee69c26d 100644 --- a/tests/slt/subquery.slt +++ b/tests/slt/subquery.slt @@ -277,14 +277,6 @@ insert into orders values (1, 1, 100), (2, 1, 200), (3, 2, 300); statement ok create index orders_user_id_index on orders(user_id); -query T -explain select id from users -where exists ( - select 1 from orders where orders.user_id = users.id -); ----- -Projection [#1] [Project => (Sort Option: Follow)] Filter #15, Is Having: false [Filter => (Sort Option: Follow)] MarkExistsApply TableScan users -> [#1, #2] [SeqScan => (Sort Option: None)] TableScan orders -> [#3, #4, #5] [IndexScan By #16 => Probe ? => (Sort Option: OrderBy: (#4 Asc Nulls Last) ignore_prefix_len: 0)] - query I rowsort select id from users where exists ( diff --git a/tests/slt/update.slt b/tests/slt/update.slt index 643540f1..13fbcc53 100644 --- a/tests/slt/update.slt +++ b/tests/slt/update.slt @@ -81,7 +81,7 @@ analyze table t_update_idx query T explain select id, b, c from t_update_idx where b = 10 ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t_update_idx -> [#1, #2, #3] [IndexScan By #2 => 10 => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t_update_idx.id, t_update_idx.b, t_update_idx.c] [Project => (Sort Option: Follow)] TableScan t_update_idx -> [t_update_idx.id, t_update_idx.b, t_update_idx.c] [IndexScan By idx_t_update_idx_b => 10 => (Sort Option: OrderBy: (t_update_idx.b Asc Nulls Last) ignore_prefix_len: 0)] statement ok update t_update_idx set c = 111 where id = 9 diff --git a/tests/slt/where_by_index_explain.slt b/tests/slt/where_by_index_explain.slt index c7ee8cdd..1a629d84 100644 --- a/tests/slt/where_by_index_explain.slt +++ b/tests/slt/where_by_index_explain.slt @@ -22,127 +22,127 @@ analyze table t1; query T explain select * from t1 limit 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3], Limit: 10 [SeqScan => (Sort Option: None)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2], Limit: 10 [SeqScan => (Sort Option: None)] query T explain select * from t1 where id = 0; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => 0 => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => 0 => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id = 0 and id = 1; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => Dummy => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => Dummy => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id = 0 and id != 0; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (#1 != 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => 0 => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (t1.id != 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => 0 => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id = 0 or id != 0 limit 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter ((#1 = 0) || (#1 != 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [SeqScan => (Sort Option: None)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter ((t1.id = 0) || (t1.id != 0)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [SeqScan => (Sort Option: None)] query T explain select * from t1 where id = 0 and id != 0 and id = 3; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (#1 != 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => Dummy => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (t1.id != 0), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => Dummy => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id = 0 and id != 0 or id = 3; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (((#1 = 0) && (#1 != 0)) || (#1 = 3)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [SeqScan => (Sort Option: None)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (((t1.id = 0) && (t1.id != 0)) || (t1.id = 3)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [SeqScan => (Sort Option: None)] query T explain select * from t1 where id > 0 and id = 3; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => 3 => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => 3 => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id >= 0 and id <= 3; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => [0, 3] => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => [0, 3] => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id <= 0 and id >= 3; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => Dummy => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => Dummy => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where (id > 10) = false; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => (-inf, 10] => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => (-inf, 10] => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where (id > 10) != true; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => (-inf, 10] => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => (-inf, 10] => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where not (id > 10); ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => (-inf, 10] => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => (-inf, 10] => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where id >= 3 or id <= 9 limit 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter ((#1 >= 3) || (#1 <= 9)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [SeqScan => (Sort Option: None)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter ((t1.id >= 3) || (t1.id <= 9)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [SeqScan => (Sort Option: None)] query T explain select * from t1 where id <= 3 or id >= 9 limit 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3], Limit: 10 [IndexScan By #0 => (-inf, 3], [9, +inf) => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2], Limit: 10 [IndexScan By pk_index => (-inf, 3], [9, +inf) => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where (id >= 0 and id <= 3) or (id >= 9 and id <= 12); ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => [0, 3], [9, 12] => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => [0, 3], [9, 12] => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where (id >= 0 or id <= 3) and (id >= 9 or id <= 12) limit 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter (((#1 >= 0) || (#1 <= 3)) && ((#1 >= 9) || (#1 <= 12))), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [SeqScan => (Sort Option: None)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Limit 10 [Limit => (Sort Option: Follow)] Filter (((t1.id >= 0) || (t1.id <= 3)) && ((t1.id >= 9) || (t1.id <= 12))), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [SeqScan => (Sort Option: None)] query T explain select * from t1 where id = 5 or (id > 5 and (id > 6 or id < 8) and id < 12); ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #0 => [5, 12) => (Sort Option: OrderBy: (#1 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By pk_index => [5, 12) => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c1 = 7 and c2 = 8; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (#3 = 8), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #1 => 7 => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (t1.c2 = 8), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By u_c1_index => 7 => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c1 = 7 and c2 < 9; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (#3 < 9), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #1 => 7 => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (t1.c2 < 9), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By u_c1_index => 7 => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where (c1 = 7 or c1 = 10) and c2 < 9; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter (#3 < 9), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #1 => 7, 10 => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter (t1.c2 < 9), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By u_c1_index => 7, 10 => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c1 is null and c2 is null; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] Filter #3 is null, Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #1 => null => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter t1.c2 is null, Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By u_c1_index => null => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c1 > 0 and c1 < 8; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #1 => (0, 8) => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By u_c1_index => (0, 8) => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c2 > 0 and c2 < 9; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #2 => (0, 9) => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By c2_index => (0, 9) => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] query T explain select * from t1 where c2 = 5; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #2 => 5 => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By c2_index => 5 => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] statement ok update t1 set c2 = 9 where c1 = 1 @@ -150,7 +150,7 @@ update t1 set c2 = 9 where c1 = 1 query T explain select * from t1 where c2 > 0 and c2 < 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #2 => (0, 10) => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By c2_index => (0, 10) => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] statement ok delete from t1 where c1 = 4 @@ -158,20 +158,20 @@ delete from t1 where c1 = 4 query T explain select * from t1 where c2 > 0 and c2 < 10; ---- -Projection [#1, #2, #3] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2, #3] [IndexScan By #2 => (0, 10) => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.id, t1.c1, t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1, t1.c2] [IndexScan By c2_index => (0, 10) => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] # unique covered query T explain select c1 from t1 where c1 < 10; ---- -Projection [#2] [Project => (Sort Option: Follow)] TableScan t1 -> [#2] [IndexScan By #1 => (-inf, 10) Covered => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c1] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.c1] [IndexScan By u_c1_index => (-inf, 10) Covered => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] # unique covered with primary key projection query T explain select c1, id from t1 where c1 < 10; ---- -Projection [#2, #1] [Project => (Sort Option: Follow)] TableScan t1 -> [#1, #2] [IndexScan By #1 => (-inf, 10) => (Sort Option: OrderBy: (#2 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c1, t1.id] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c1] [IndexScan By u_c1_index => (-inf, 10) => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last) ignore_prefix_len: 0)] statement ok drop index t1.u_c1_index; @@ -180,7 +180,7 @@ drop index t1.u_c1_index; query T explain select c2 from t1 where c2 < 10 and c2 > 0; ---- -Projection [#3] [Project => (Sort Option: Follow)] TableScan t1 -> [#3] [IndexScan By #2 => (0, 10) Covered => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c2] [Project => (Sort Option: Follow)] TableScan t1 -> [t1.c2] [IndexScan By c2_index => (0, 10) Covered => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] statement ok insert into t1 values(100000002, 100000002, 8); @@ -189,7 +189,7 @@ insert into t1 values(100000002, 100000002, 8); query T explain select distinct c2 from t1 where c2 < 10 and c2 > 0; ---- -Projection [#3] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [#3] [StreamDistinct => (Sort Option: Follow)] TableScan t1 -> [#3] [IndexScan By #2 => (0, 10) Covered => (Sort Option: OrderBy: (#3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c2] [Project => (Sort Option: Follow)] Aggregate [] -> Group By [t1.c2] [StreamDistinct => (Sort Option: Follow)] TableScan t1 -> [t1.c2] [IndexScan By c2_index => (0, 10) Covered => (Sort Option: OrderBy: (t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] statement ok delete from t1 where id = 100000002; @@ -201,13 +201,13 @@ drop index t1.c2_index; query T explain select c1, c2 from t1 where c1 < 10 and c1 > 0 and c2 >0 and c2 < 10; ---- -Projection [#2, #3] [Project => (Sort Option: Follow)] Filter ((#3 > 0) && (#3 < 10)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#2, #3] [IndexScan By #3 => ((0), (10)) Covered => (Sort Option: OrderBy: (#2 Asc Nulls Last, #3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c1, t1.c2] [Project => (Sort Option: Follow)] Filter ((t1.c2 > 0) && (t1.c2 < 10)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.c1, t1.c2] [IndexScan By p_index => ((0), (10)) Covered => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last, t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] # composite covered projection reorder query T explain select c2, c1 from t1 where c1 < 10 and c1 > 0 and c2 > 0 and c2 < 10; ---- -Projection [#3, #2] [Project => (Sort Option: Follow)] Filter ((#3 > 0) && (#3 < 10)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [#2, #3] [IndexScan By #3 => ((0), (10)) Covered => (Sort Option: OrderBy: (#2 Asc Nulls Last, #3 Asc Nulls Last) ignore_prefix_len: 0)] +Projection [t1.c2, t1.c1] [Project => (Sort Option: Follow)] Filter ((t1.c2 > 0) && (t1.c2 < 10)), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.c1, t1.c2] [IndexScan By p_index => ((0), (10)) Covered => (Sort Option: OrderBy: (t1.c1 Asc Nulls Last, t1.c2 Asc Nulls Last) ignore_prefix_len: 0)] statement ok @@ -230,7 +230,7 @@ create index idx_cover on t_cover (c1, c2, c3); query T explain select c2, c3 from t_cover where c1 = 2; ---- -Projection [#3, #4] [Project => (Sort Option: Follow)] Filter (#2 = 2), Is Having: false [Filter => (Sort Option: Follow)] TableScan t_cover -> [#2, #3, #4] [SeqScan => (Sort Option: None)] +Projection [t_cover.c2, t_cover.c3] [Project => (Sort Option: Follow)] Filter (t_cover.c1 = 2), Is Having: false [Filter => (Sort Option: Follow)] TableScan t_cover -> [t_cover.c1, t_cover.c2, t_cover.c3] [SeqScan => (Sort Option: None)] statement ok drop table t_cover; diff --git a/tests/slt/window.slt b/tests/slt/window.slt index 3d3ba6f9..54309212 100644 --- a/tests/slt/window.slt +++ b/tests/slt/window.slt @@ -45,7 +45,7 @@ explain select row_number() over (order by v desc, id) from window_test ---- -Projection [#4, #5, #6] [Project => (Sort Option: Follow)] Window [WindowFunction { kind: RowNumber, args: [], ty: Bigint }] -> Order By [#3 Desc Nulls Last, #1 Asc Nulls Last] [Window => (Sort Option: OrderBy: (#3 Desc Nulls Last, #1 Asc Nulls Last) ignore_prefix_len: 0)] Sort By #3 Desc Nulls Last, #1 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#3 Desc Nulls Last, #1 Asc Nulls Last) ignore_prefix_len: 0)] Window [WindowFunction { kind: RowNumber, args: [], ty: Bigint }, WindowFunction { kind: Rank, args: [], ty: Bigint }] -> Partition By [#2] Order By [#3 Asc Nulls Last, #1 Asc Nulls Last] [Window => (Sort Option: OrderBy: (#2 Asc Nulls Last, #3 Asc Nulls Last, #1 Asc Nulls Last) ignore_prefix_len: 0)] Sort By #2 Asc Nulls Last, #3 Asc Nulls Last, #1 Asc Nulls Last [Sort => (Sort Option: OrderBy: (#2 Asc Nulls Last, #3 Asc Nulls Last, #1 Asc Nulls Last) ignore_prefix_len: 0)] TableScan window_test -> [#1, #2, #3] [SeqScan => (Sort Option: None)] +Projection [row_number() over (partition by window_test.k order by window_test.v Asc Nulls Last, window_test.id Asc Nulls Last), rank() over (partition by window_test.k order by window_test.v Asc Nulls Last, window_test.id Asc Nulls Last), row_number() over (order by window_test.v Desc Nulls Last, window_test.id Asc Nulls Last)] [Project => (Sort Option: Follow)] Window [WindowFunction { kind: RowNumber, args: [], ty: Bigint }] -> Order By [window_test.v Desc Nulls Last, window_test.id Asc Nulls Last] [Window => (Sort Option: OrderBy: (window_test.v Desc Nulls Last, window_test.id Asc Nulls Last) ignore_prefix_len: 0)] Sort By window_test.v Desc Nulls Last, window_test.id Asc Nulls Last [Sort => (Sort Option: OrderBy: (window_test.v Desc Nulls Last, window_test.id Asc Nulls Last) ignore_prefix_len: 0)] Window [WindowFunction { kind: Rank, args: [], ty: Bigint }] -> Partition By [window_test.k] Order By [window_test.v Asc Nulls Last, window_test.id Asc Nulls Last] [Window => (Sort Option: OrderBy: (window_test.k Asc Nulls Last, window_test.v Asc Nulls Last, window_test.id Asc Nulls Last) ignore_prefix_len: 0)] Window [WindowFunction { kind: RowNumber, args: [], ty: Bigint }] -> Partition By [window_test.k] Order By [window_test.v Asc Nulls Last, window_test.id Asc Nulls Last] [Window => (Sort Option: OrderBy: (window_test.k Asc Nulls Last, window_test.v Asc Nulls Last, window_test.id Asc Nulls Last) ignore_prefix_len: 0)] Sort By window_test.k Asc Nulls Last, window_test.v Asc Nulls Last, window_test.id Asc Nulls Last [Sort => (Sort Option: OrderBy: (window_test.k Asc Nulls Last, window_test.v Asc Nulls Last, window_test.id Asc Nulls Last) ignore_prefix_len: 0)] TableScan window_test -> [window_test.id, window_test.k, window_test.v] [SeqScan => (Sort Option: None)] query III select id, From 113fb7d8a131a811c25f188fce4296089f55f6c8 Mon Sep 17 00:00:00 2001 From: kould Date: Thu, 3 Sep 2026 05:17:01 +0800 Subject: [PATCH 2/2] fix: preserve arena expression ownership boundaries --- kite_sql_serde_macros/src/orm.rs | 86 ++++++++------------- src/binder/aggregate.rs | 21 +++-- src/binder/create_index.rs | 2 +- src/binder/expr.rs | 20 +++-- src/binder/select.rs | 6 +- src/catalog/table.rs | 1 + src/expression/mod.rs | 9 +++ src/expression/visitor_mut.rs | 13 ++++ src/orm/ddl.rs | 128 +++++++++++++++++-------------- src/orm/mod.rs | 100 ++++++++++++++---------- src/planner/arena.rs | 27 +++++-- src/planner/mod.rs | 28 +++++++ tests/macros-test/src/main.rs | 2 +- 13 files changed, 267 insertions(+), 176 deletions(-) diff --git a/kite_sql_serde_macros/src/orm.rs b/kite_sql_serde_macros/src/orm.rs index ad0f004c..24d53f26 100644 --- a/kite_sql_serde_macros/src/orm.rs +++ b/kite_sql_serde_macros/src/orm.rs @@ -71,18 +71,14 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { let mut field_index_resolvers = Vec::new(); let mut params = Vec::new(); let mut orm_fields = Vec::new(); - let mut orm_columns = Vec::new(); let mut field_getters = Vec::new(); let mut column_names = Vec::new(); - let mut placeholder_names = Vec::new(); let mut orm_indexes = Vec::new(); let mut persisted_columns = Vec::new(); let mut index_names = BTreeSet::new(); index_names.insert("pk_index".to_string()); let mut primary_key_type = None; let mut primary_key_value = None; - let mut primary_key_column = None; - let mut primary_key_placeholder = None; let mut primary_key_count = 0usize; for field in data_struct.fields { @@ -165,23 +161,18 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { .map(|value| LitStr::new(&value, Span::call_site())); let field_name_string = field_name.to_string(); let column_name = field.rename.unwrap_or_else(|| field_name_string.clone()); - let placeholder_name = format!(":{column_name}"); let column_name_lit = LitStr::new(&column_name, Span::call_site()); - let placeholder_lit = LitStr::new(&placeholder_name, Span::call_site()); let is_primary_key = field.primary_key; let is_unique = field.unique; let is_index = field.index; - let column_index = orm_columns.len(); + let column_index = orm_fields.len(); let field_index_ident = format_ident!("__kite_orm_{field_name}_index"); persisted_columns.push((field_name_string, column_name.clone())); column_names.push(column_name.clone()); - placeholder_names.push(placeholder_name.clone()); if is_primary_key { primary_key_count += 1; - primary_key_column = Some(column_name.clone()); - primary_key_placeholder = Some(placeholder_name.clone()); primary_key_type = Some(quote! { #field_ty }); primary_key_value = Some(quote! { &self.#field_name @@ -250,52 +241,40 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { )? }); params.push(quote! { - (#placeholder_lit, ::kite_sql::orm::ToDataValue::to_data_value(&self.#field_name)) + (#column_name_lit, ::kite_sql::orm::ToDataValue::to_data_value(&self.#field_name)) }); orm_fields.push(quote! { - ::kite_sql::orm::OrmField { - column: #column_name_lit, - column_index: #column_index, - placeholder: #placeholder_lit, - primary_key: #is_primary_key, - unique: #is_unique, - } - }); - let getter_name = format_ident!("{}", field_name); - field_getters.push(quote! { - pub fn #getter_name() -> ::kite_sql::orm::Field { - ::kite_sql::orm::Field::new(#table_name_lit, #column_name_lit) - } - }); - orm_columns.push(quote! { { let data_type = #data_type; - let default = #default_tokens - .map(|value| { - arena.alloc_expression(::kite_sql::expression::ScalarExpression::Constant( - value - .cast(&data_type) - .expect("failed to cast ORM default value to column type"), - )) - }); - let desc = ::kite_sql::catalog::column::ColumnDesc::new( + let default = #default_tokens.map(|value| { + ::kite_sql::expression::ScalarExpression::Constant( + value + .cast(&data_type) + .expect("failed to cast ORM default value to column type"), + ) + }); + ::kite_sql::orm::OrmField { + column: #column_name_lit, + column_index: #column_index, data_type, - #is_primary_key.then_some(#column_index), - #is_unique, - default, - ) - .expect("failed to build ORM column descriptor"); - ::kite_sql::catalog::column::ColumnCatalog::new( - #column_name_lit.to_string(), - if #is_primary_key { + nullable: if #is_primary_key { false } else { <#field_ty as ::kite_sql::orm::ModelColumnType>::nullable() }, - desc, - ) + default, + primary_key: #is_primary_key, + unique: #is_unique, + } } }); + let getter_name = format_ident!("{}", field_name); + field_getters.push(quote! { + pub fn #getter_name() -> ::kite_sql::orm::Field { + ::kite_sql::orm::Field::new(#table_name_lit, #column_name_lit) + } + }); + if is_unique { let unique_index_name_value = format!("uk_{column_name}_index"); if !index_names.insert(unique_index_name_value.clone()) { @@ -398,8 +377,6 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { let primary_key_type = primary_key_type.expect("primary key checked above"); let primary_key_value = primary_key_value.expect("primary key checked above"); - let _primary_key_column = primary_key_column.expect("primary key checked above"); - let _primary_key_placeholder = primary_key_placeholder.expect("primary key checked above"); let field_count = field_index_declarations.len(); let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); @@ -444,15 +421,12 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { } fn fields() -> &'static [::kite_sql::orm::OrmField] { - &[ - #(#orm_fields),* - ] - } - - fn columns(arena: &mut ::kite_sql::planner::TableArena) -> ::std::vec::Vec<::kite_sql::catalog::column::ColumnCatalog> { - vec![ - #(#orm_columns),* - ] + static ORM_FIELDS: ::std::sync::LazyLock<::std::vec::Vec<::kite_sql::orm::OrmField>> = ::std::sync::LazyLock::new(|| { + vec![ + #(#orm_fields),* + ] + }); + ORM_FIELDS.as_slice() } fn indexes() -> &'static [(&'static str, &'static [&'static str], bool)] { diff --git a/src/binder/aggregate.rs b/src/binder/aggregate.rs index 254371d6..c4226bd4 100644 --- a/src/binder/aggregate.rs +++ b/src/binder/aggregate.rs @@ -72,12 +72,12 @@ impl> Binder<'_, '_, T, A> &mut self, select_list: &mut [ExprRef], mut group_by_exprs: Vec, - arena: &PlanArena<'_>, + arena: &mut PlanArena<'_>, ) -> Result<(), DatabaseError> { self.validate_groupby_illegal_column(select_list, &group_by_exprs, arena)?; for expr in group_by_exprs.iter_mut() { - self.visit_group_by_expr(select_list, *expr, arena); + self.visit_group_by_expr(select_list, *expr, arena)?; } Ok(()) } @@ -214,8 +214,8 @@ impl> Binder<'_, '_, T, A> &mut self, select_list: &mut [ExprRef], expr: ExprRef, - arena: &PlanArena<'_>, - ) { + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { alias, .. } = arena.expression(expr) { if let Some(i) = select_list.iter().position(|inner_expr| { if let ScalarExpression::Alias { @@ -227,8 +227,12 @@ impl> Binder<'_, '_, T, A> false } }) { - self.context.group_by_exprs.push(select_list[i]); - return; + // GROUP BY evaluates against the aggregate input, while the select + // expression is later rewritten against aggregate output. + self.context + .group_by_exprs + .push(select_list[i].clone_expression(arena)?); + return Ok(()); } } @@ -236,8 +240,11 @@ impl> Binder<'_, '_, T, A> .iter() .position(|column| column.eq_ignore_colref_pos(expr, arena)) { - self.context.group_by_exprs.push(select_list[i]) + self.context + .group_by_exprs + .push(select_list[i].clone_expression(arena)?); } + Ok(()) } /// Validate having or orderby clause is valid, if SQL has group by clause. diff --git a/src/binder/create_index.rs b/src/binder/create_index.rs index 44187aec..621f48a0 100644 --- a/src/binder/create_index.rs +++ b/src/binder/create_index.rs @@ -37,7 +37,7 @@ impl> Binder<'_, '_, T, A> Source::Table(table) => { TableScanOperator::build(table_name.clone(), table, true, arena) } - Source::View(view) => Ok(LogicalPlan::clone(&view.plan)), + Source::View(view) => view.plan.clone_plan(arena), Source::Schema(_) => Err(DatabaseError::UnsupportedStmt( "derived source cannot be rebound as a base relation".to_string(), )), diff --git a/src/binder/expr.rs b/src/binder/expr.rs index 5470c1ed..a0e487ec 100644 --- a/src/binder/expr.rs +++ b/src/binder/expr.rs @@ -90,14 +90,17 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T for (position, expr) in exprs.into_iter().enumerate() { let (alias_expr, alias_ref) = self.bind_temp_table_alias(expr, position, arena); + let predicate_alias = arena.alloc_expression(alias_expr); + // The projection evaluates against the subquery-local tuple, while the + // predicate evaluates against the combined left/right tuple. Clone the + // complete expression graph so position rewrites in either context do + // not mutate the other one through shared ExprRefs. + let projection_alias = predicate_alias.clone_expression(arena)?; if !is_tuple { - let alias_plan = Self::build_project_plan( - sub_query, - vec![arena.alloc_expression(alias_expr.clone())], - ); - return Ok((alias_expr, alias_plan)); + let alias_plan = Self::build_project_plan(sub_query, vec![projection_alias]); + return Ok((arena.expression(predicate_alias).clone(), alias_plan)); } - alias_exprs.push(arena.alloc_expression(alias_expr)); + alias_exprs.push(projection_alias); alias_refs.push(alias_ref); } @@ -329,8 +332,11 @@ impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T .iter() .find(|((table, column), _)| table.is_none() && column == column_name) { + // ORDER BY evaluates the alias against projection output, while + // the select expression is rewritten against projection input. + // Keep their position-sensitive expression graphs independent. return Ok(ScalarExpression::Alias { - expr: *expr, + expr: expr.clone_expression(arena)?, alias: AliasType::Name(column_name.to_string()), }); } diff --git a/src/binder/select.rs b/src/binder/select.rs index ac358302..da6875e1 100644 --- a/src/binder/select.rs +++ b/src/binder/select.rs @@ -1201,7 +1201,11 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' Source::Table(table) => { TableScanOperator::build(table_name.clone(), table, with_pk, arena)? } - Source::View(view) => LogicalPlan::clone(&view.plan), + Source::View(view) => { + // Cached view expressions live in the persistent arena. Clone the + // complete graph before optimizer passes rewrite expression nodes. + view.plan.clone_plan(arena)? + } Source::Schema(_) => { return Err(DatabaseError::UnsupportedStmt( "derived source cannot be rebound as a base relation".to_string(), diff --git a/src/catalog/table.rs b/src/catalog/table.rs index 27ce3cbe..9f8eb0f1 100644 --- a/src/catalog/table.rs +++ b/src/catalog/table.rs @@ -321,6 +321,7 @@ impl TableCatalog { .map(|index| source_arena.index(*index).clone()) .collect_vec(); + source_arena.materialize_expressions_into_table_arena(); Self::reload( self.name.clone(), column_catalogs.into_iter(), diff --git a/src/expression/mod.rs b/src/expression/mod.rs index 29d5ab0d..3e79bb45 100644 --- a/src/expression/mod.rs +++ b/src/expression/mod.rs @@ -735,6 +735,15 @@ impl ExprRef { eq_col::eq_ignore_colref_pos(self, other, arena) } + pub(crate) fn clone_expression( + self, + arena: &mut PlanArena<'_>, + ) -> Result { + let mut cloned = self; + crate::expression::visitor_mut::ExprCloner.visit(&mut cloned, arena)?; + Ok(cloned) + } + pub fn unpack_alias(self, arena: &impl MetaArena) -> ExprRef { if let ScalarExpression::Alias { alias: AliasType::Expr(expr), diff --git a/src/expression/visitor_mut.rs b/src/expression/visitor_mut.rs index d99978eb..4eea5c5a 100644 --- a/src/expression/visitor_mut.rs +++ b/src/expression/visitor_mut.rs @@ -26,6 +26,19 @@ use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluat use crate::types::value::DataValue; use crate::types::LogicalType; +pub(crate) struct ExprCloner; + +impl ExprVisitorMut for ExprCloner { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + *expr = arena.alloc_expression(arena.expression(*expr).clone()); + walk_mut_expr(self, expr, arena) + } +} + pub(crate) struct PositionShift { pub(crate) delta: isize, } diff --git a/src/orm/ddl.rs b/src/orm/ddl.rs index e213e7b3..322227bb 100644 --- a/src/orm/ddl.rs +++ b/src/orm/ddl.rs @@ -1,5 +1,11 @@ use super::*; +enum OrmDefaultChange { + NoChange, + Set(ScalarExpression), + Drop, +} + impl Database { fn table_catalog(&self, table_name: &str) -> Result, DatabaseError> { let transaction = self.storage.transaction()?; @@ -140,10 +146,10 @@ impl Database { /// when the underlying DDL supports them. Primary-key changes and unique /// constraint changes still return an error so you can handle them manually. pub fn migrate(&mut self) -> Result<(), DatabaseError> { - let columns = M::columns(self.state.table_arena().borrow_mut()); - if columns.is_empty() { + let fields = M::fields(); + if fields.is_empty() { return Err(DatabaseError::UnsupportedStmt( - "ORM migration requires Model::columns(); #[derive(Model)] provides it automatically" + "ORM migration requires Model::fields(); #[derive(Model)] provides it automatically" .to_string(), )); } @@ -168,11 +174,11 @@ impl Database { (table_primary_key, current_columns) }; - let model_primary_key = columns + let model_primary_key = fields .iter() - .find(|column| column.desc().is_primary()) + .find(|field| field.primary_key) .ok_or(DatabaseError::PrimaryKeyNotFound)?; - if table_primary_key.name() != model_primary_key.name() + if table_primary_key.name() != model_primary_key.column || !model_column_matches_catalog( model_primary_key, &table_primary_key, @@ -184,80 +190,78 @@ impl Database { M::table_name(), ))); } - let model_columns = columns + let model_columns = fields .iter() - .map(|column| (column.name(), column)) + .map(|field| (field.column, field)) .collect::>(); let mut handled_current = BTreeMap::new(); let mut handled_model = BTreeMap::new(); - for column in &columns { - let Some(current_column) = current_columns.get(column.name()) else { + for field in fields { + let Some(current_column) = current_columns.get(field.column) else { continue; }; handled_current.insert(current_column.name().to_string(), ()); - handled_model.insert(column.name(), ()); + handled_model.insert(field.column, ()); - if column.desc().is_primary() != current_column.desc().is_primary() { + if field.primary_key != current_column.desc().is_primary() { return Err(DatabaseError::InvalidValue(::std::format!( "ORM migration does not support changing the primary key for table `{}`", M::table_name(), ))); } - if column.desc().is_unique() != current_column.desc().is_unique() { + if field.unique != current_column.desc().is_unique() { return Err(DatabaseError::InvalidValue(::std::format!( "ORM migration cannot automatically change unique constraint on column `{}` of table `{}`", - column.name(), + field.column, M::table_name(), ))); } if model_column_matches_catalog( - column, + field, current_column, &PlanArena::new(self.state.table_arena()), )? { continue; } - if !model_column_type_matches_catalog(column, current_column) { + if !model_column_type_matches_catalog(field, current_column) { execute_change_column( self, M::table_name(), - column.name(), - column.name(), - column.datatype().clone(), - DefaultChange::NoChange, + field.column, + field.column, + field.data_type.clone(), + OrmDefaultChange::NoChange, NotNullChange::NoChange, )?; } let arena = PlanArena::new(self.state.table_arena()); - if model_column_default(column, &arena)? - != catalog_column_default(current_column, &arena)? - { + if model_column_default(field) != catalog_column_default(current_column, &arena) { execute_change_column( self, M::table_name(), - column.name(), - column.name(), - column.datatype().clone(), - match column.desc().default { - Some(expr) => DefaultChange::Set(expr), - None => DefaultChange::Drop, + field.column, + field.column, + field.data_type.clone(), + match field.default.clone() { + Some(expr) => OrmDefaultChange::Set(expr), + None => OrmDefaultChange::Drop, }, NotNullChange::NoChange, )?; } - if column.nullable() != current_column.nullable() { + if field.nullable != current_column.nullable() { execute_change_column( self, M::table_name(), - column.name(), - column.name(), - column.datatype().clone(), - DefaultChange::NoChange, - if column.nullable() { + field.column, + field.column, + field.data_type.clone(), + OrmDefaultChange::NoChange, + if field.nullable { NotNullChange::Drop } else { NotNullChange::Set @@ -267,9 +271,9 @@ impl Database { } let mut rename_pairs = Vec::new(); - let unmatched_model_columns = columns + let unmatched_model_columns = fields .iter() - .filter(|column| !handled_model.contains_key(column.name())) + .filter(|field| !handled_model.contains_key(field.column)) .collect::>(); let unmatched_current_columns = current_columns .values() @@ -277,8 +281,8 @@ impl Database { .cloned() .collect::>(); - for model_column in &unmatched_model_columns { - if model_column.desc().is_primary() { + for model_field in &unmatched_model_columns { + if model_field.primary_key { continue; } let mut candidates = Vec::new(); @@ -287,7 +291,7 @@ impl Database { .filter(|column| !column.desc().is_primary()) { if model_column_rename_compatible( - model_column, + model_field, column, &PlanArena::new(self.state.table_arena()), )? { @@ -301,7 +305,7 @@ impl Database { let mut reverse_candidates = Vec::new(); for other in unmatched_model_columns .iter() - .filter(|other| !other.desc().is_primary()) + .filter(|other| !other.primary_key) { if model_column_rename_compatible( other, @@ -314,9 +318,9 @@ impl Database { if reverse_candidates.len() != 1 { continue; } - rename_pairs.push((current_column.name().to_string(), model_column.name())); + rename_pairs.push((current_column.name().to_string(), model_field.column)); handled_current.insert(current_column.name().to_string(), ()); - handled_model.insert(model_column.name(), ()); + handled_model.insert(model_field.column, ()); } for (old_name, new_name) in rename_pairs { @@ -329,7 +333,7 @@ impl Database { &old_name, new_name, current_column.datatype().clone(), - DefaultChange::NoChange, + OrmDefaultChange::NoChange, NotNullChange::NoChange, )?; } @@ -351,21 +355,21 @@ impl Database { execute_drop_column(self, M::table_name(), column.name())?; } - for column in &columns { - if handled_model.contains_key(column.name()) - || current_columns.contains_key(column.name()) + for field in fields { + if handled_model.contains_key(field.column) + || current_columns.contains_key(field.column) { continue; } - if column.desc().is_primary() { + if field.primary_key { return Err(DatabaseError::InvalidValue(::std::format!( "ORM migration cannot add a new primary key column `{}` to an existing table `{}`", - column.name(), + field.column, M::table_name(), ))); } - execute_add_column(self, M::table_name(), column)?; + execute_add_column(self, M::table_name(), field)?; } for index in M::indexes() { @@ -428,8 +432,11 @@ fn execute_create_table( database: &mut Database, if_not_exists: bool, ) -> Result<(), DatabaseError> { - let columns = M::columns(database.state.table_arena().borrow_mut()); - database.execute_mut("ORM CREATE TABLE", &[], move |binder, _| { + database.execute_mut("ORM CREATE TABLE", &[], move |binder, arena| { + let columns = M::fields() + .iter() + .map(|field| field.to_column_catalog(arena)) + .collect::, _>>()?; binder.bind_create_table(M::table_name().into(), columns, if_not_exists) }) } @@ -517,12 +524,17 @@ fn execute_change_column( old_column_name: &str, new_column_name: &str, data_type: LogicalType, - default_change: DefaultChange, + default_change: OrmDefaultChange, not_null_change: NotNullChange, ) -> Result<(), DatabaseError> { let old_column_name = old_column_name.to_string(); let new_column_name = new_column_name.to_string(); - database.execute_mut("ORM CHANGE COLUMN", &[], move |binder, _| { + database.execute_mut("ORM CHANGE COLUMN", &[], move |binder, arena| { + let default_change = match default_change { + OrmDefaultChange::Set(expr) => DefaultChange::Set(arena.alloc_expression(expr)), + OrmDefaultChange::Drop => DefaultChange::Drop, + OrmDefaultChange::NoChange => DefaultChange::NoChange, + }; binder.bind_change_column( table_name.into(), old_column_name, @@ -548,11 +560,11 @@ fn execute_drop_column( fn execute_add_column( database: &mut Database, table_name: &'static str, - column: &ColumnCatalog, + field: &OrmField, ) -> Result<(), DatabaseError> { - let column = column.clone(); - database.execute_mut("ORM ADD COLUMN", &[], move |binder, _| { - binder.bind_add_column(table_name.into(), column, false) + let field = field.clone(); + database.execute_mut("ORM ADD COLUMN", &[], move |binder, arena| { + binder.bind_add_column(table_name.into(), field.to_column_catalog(arena)?, false) }) } diff --git a/src/orm/mod.rs b/src/orm/mod.rs index f1249f64..2d8958e5 100644 --- a/src/orm/mod.rs +++ b/src/orm/mod.rs @@ -37,7 +37,7 @@ mod ddl; mod dml; mod dql; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq)] /// Static metadata about a single model field. /// /// This type is primarily consumed by code generated from `#[derive(Model)]`. @@ -45,11 +45,32 @@ mod dql; pub struct OrmField { pub column: &'static str, pub column_index: usize, - pub placeholder: &'static str, + pub data_type: LogicalType, + pub nullable: bool, + pub default: Option, pub primary_key: bool, pub unique: bool, } +impl OrmField { + fn to_column_catalog(&self, arena: &mut PlanArena<'_>) -> Result { + let default = self + .default + .clone() + .map(|expr| arena.alloc_expression(expr)); + Ok(ColumnCatalog::new( + self.column.to_string(), + self.nullable, + crate::catalog::ColumnDesc::new( + self.data_type.clone(), + self.primary_key.then_some(self.column_index), + self.unique, + default, + )?, + )) + } +} + /// One row returned by [`Database::describe`] or [`DBTransaction::describe`]. #[derive(Debug, Clone, PartialEq, Eq)] pub struct DescribeColumn { @@ -2659,8 +2680,8 @@ where .ok_or_else(|| DatabaseError::column_not_found(field.column.to_string()))?; let column_catalog = arena.column(column); let value = params - .get(field.placeholder) - .ok_or_else(|| DatabaseError::parameter_not_found(field.placeholder))? + .get(field.column) + .ok_or_else(|| DatabaseError::parameter_not_found(field.column))? .clone() .cast(column_catalog.datatype())?; value.check_len(column_catalog.datatype())?; @@ -2704,14 +2725,6 @@ pub trait Model: Sized + FromQueryRow { /// Returns metadata for every persisted field on the model. fn fields() -> &'static [OrmField]; - /// Returns persisted column catalogs for the model. - /// - /// `#[derive(Model)]` generates this automatically. Manual implementations - /// can override it to opt into [`Database::migrate`](crate::orm::Database::migrate). - fn columns(_arena: &mut crate::planner::TableArena) -> Vec { - Vec::new() - } - /// Returns secondary indexes declared by the model. fn indexes() -> &'static [(&'static str, &'static [&'static str], bool)] { &[] @@ -3210,46 +3223,43 @@ impl_from_query_tuple!( (A, B, C, D, E, F, G, H), ); -fn model_column_default( - model: &ColumnCatalog, - arena: &PlanArena<'_>, -) -> Result, DatabaseError> { - model.default_value(arena) +fn model_column_default(model: &OrmField) -> Option<&ScalarExpression> { + model.default.as_ref() } -fn catalog_column_default( +fn catalog_column_default<'a>( column: &ColumnCatalog, - arena: &PlanArena<'_>, -) -> Result, DatabaseError> { - column.default_value(arena) + arena: &'a PlanArena<'_>, +) -> Option<&'a ScalarExpression> { + column.desc().default.map(|expr| arena.expression(expr)) } -fn model_column_type_matches_catalog(model: &ColumnCatalog, column: &ColumnCatalog) -> bool { - model.datatype() == column.datatype() +fn model_column_type_matches_catalog(model: &OrmField, column: &ColumnCatalog) -> bool { + model.data_type == *column.datatype() } fn model_column_matches_catalog( - model: &ColumnCatalog, + model: &OrmField, column: &ColumnCatalog, arena: &PlanArena<'_>, ) -> Result { - Ok(model.desc().is_primary() == column.desc().is_primary() - && model.desc().is_unique() == column.desc().is_unique() - && model.nullable() == column.nullable() + Ok(model.primary_key == column.desc().is_primary() + && model.unique == column.desc().is_unique() + && model.nullable == column.nullable() && model_column_type_matches_catalog(model, column) - && model_column_default(model, arena)? == catalog_column_default(column, arena)?) + && model_column_default(model) == catalog_column_default(column, arena)) } fn model_column_rename_compatible( - model: &ColumnCatalog, + model: &OrmField, column: &ColumnCatalog, arena: &PlanArena<'_>, ) -> Result { - Ok(model.desc().is_primary() == column.desc().is_primary() - && model.desc().is_unique() == column.desc().is_unique() - && model.nullable() == column.nullable() + Ok(model.primary_key == column.desc().is_primary() + && model.unique == column.desc().is_unique() + && model.nullable == column.nullable() && model_column_type_matches_catalog(model, column) - && model_column_default(model, arena)? == catalog_column_default(column, arena)?) + && model_column_default(model) == catalog_column_default(column, arena)) } fn extract_optional_model(iter: I) -> Result, DatabaseError> @@ -3392,21 +3402,27 @@ mod tests { OrmField { column: "id", column_index: 0, - placeholder: "id", + data_type: LogicalType::Integer, + nullable: false, + default: None, primary_key: true, unique: false, }, OrmField { column: "name", column_index: 1, - placeholder: "name", + data_type: LogicalType::Varchar(None, CharLengthUnits::Characters), + nullable: false, + default: None, primary_key: false, unique: false, }, OrmField { column: "age", column_index: 2, - placeholder: "age", + data_type: LogicalType::Integer, + nullable: true, + default: None, primary_key: false, unique: false, }, @@ -3474,21 +3490,27 @@ mod tests { OrmField { column: "id", column_index: 0, - placeholder: "id", + data_type: LogicalType::Integer, + nullable: false, + default: None, primary_key: true, unique: false, }, OrmField { column: "user_id", column_index: 1, - placeholder: "user_id", + data_type: LogicalType::Integer, + nullable: false, + default: None, primary_key: false, unique: false, }, OrmField { column: "amount", column_index: 2, - placeholder: "amount", + data_type: LogicalType::Integer, + nullable: false, + default: None, primary_key: false, unique: false, }, diff --git a/src/planner/arena.rs b/src/planner/arena.rs index c17a709f..dafd9e63 100644 --- a/src/planner/arena.rs +++ b/src/planner/arena.rs @@ -427,6 +427,26 @@ impl<'a> PlanArena<'a> { self.table_arena } + fn append_expressions_to_table_arena(&self, table_arena: &mut TableArena) { + for expression in &self.expressions { + table_arena.expressions.push(TableArenaExpression { + expression: expression.clone(), + live: true, + }); + } + } + + pub(crate) fn materialize_expressions_into_table_arena(&self) { + self.assert_table_arena_unchanged(); + if self.expressions.is_empty() { + return; + } + + let table_arena = self.table_arena.borrow_mut(); + self.append_expressions_to_table_arena(table_arena); + table_arena.increment_version(); + } + pub(crate) fn materialize_into_table_arena(&self) { self.assert_table_arena_unchanged(); @@ -451,12 +471,7 @@ impl<'a> PlanArena<'a> { live: true, }); } - for expression in &self.expressions { - table_arena.expressions.push(TableArenaExpression { - expression: expression.clone(), - live: true, - }); - } + self.append_expressions_to_table_arena(table_arena); if !self.columns.is_empty() || !self.indexes.is_empty() || !self.expressions.is_empty() { table_arena.increment_version(); } diff --git a/src/planner/mod.rs b/src/planner/mod.rs index a80b0571..93a2b979 100644 --- a/src/planner/mod.rs +++ b/src/planner/mod.rs @@ -17,10 +17,12 @@ pub mod operator; use crate::catalog::TableName; use crate::errors::DatabaseError; +use crate::expression::visitor_mut::ExprCloner; use crate::planner::operator::recursive_cte::{RecursiveCteOperator, RecursiveScanOperator}; use crate::planner::operator::set_membership::SetMembershipOperator; use crate::planner::operator::union::UnionOperator; use crate::planner::operator::values::ValuesOperator; +use crate::planner::operator::visitor_mut::{OperatorExprVisitorMut, OperatorVisitorMut}; use crate::planner::operator::{Operator, PhysicalOption}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt; @@ -150,6 +152,32 @@ impl LogicalPlan { } } + pub(crate) fn clone_plan( + &self, + arena: &mut PlanArena<'_>, + ) -> Result { + fn clone_expressions( + plan: &mut LogicalPlan, + cloner: &mut ExprCloner, + arena: &mut PlanArena<'_>, + ) -> Result<(), DatabaseError> { + OperatorExprVisitorMut::new(cloner, arena).visit_operator(&mut plan.operator)?; + match plan.childrens.as_mut() { + Childrens::Only(child) => clone_expressions(child, cloner, arena)?, + Childrens::Twins { left, right } => { + clone_expressions(left, cloner, arena)?; + clone_expressions(right, cloner, arena)?; + } + Childrens::None => {} + } + Ok(()) + } + + let mut plan = self.clone(); + clone_expressions(&mut plan, &mut ExprCloner, arena)?; + Ok(plan) + } + pub(crate) fn take(&mut self) -> Self { std::mem::replace(self, Self::new(Operator::Dummy, Childrens::None)) } diff --git a/tests/macros-test/src/main.rs b/tests/macros-test/src/main.rs index 5d9d1746..5aff4964 100644 --- a/tests/macros-test/src/main.rs +++ b/tests/macros-test/src/main.rs @@ -2948,7 +2948,7 @@ mod test { })?; assert_eq!( plan, - "Projection [users.user_name] [Project => (Sort Option: Follow)] TableScan users -> [users.id, users.user_name] [IndexScan By pk_index => 1 => (Sort Option: OrderBy: (users.id Asc Nulls Last) ignore_prefix_len: 0)]" + "Projection [users.user_name] [Project => (Sort Option: Follow)] Filter (users.id = 1), Is Having: false [Filter => (Sort Option: Follow)] TableScan users -> [users.id, users.user_name] [SeqScan => (Sort Option: None)]" ); let set_plan = database.explain(|ctx| {