diff --git a/kite_sql_serde_macros/src/projection.rs b/kite_sql_serde_macros/src/projection.rs index 9c64c4cf..fff67210 100644 --- a/kite_sql_serde_macros/src/projection.rs +++ b/kite_sql_serde_macros/src/projection.rs @@ -108,7 +108,7 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { ) -> ::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)]>, + A: AsRef<[(usize, ::kite_sql::types::LogicalType)]>, { Ok(::std::vec![ #(#projection_exprs.into_scalar()),* diff --git a/kite_sql_serde_macros/src/reference_serialization.rs b/kite_sql_serde_macros/src/reference_serialization.rs index 3a46a193..f8d11644 100644 --- a/kite_sql_serde_macros/src/reference_serialization.rs +++ b/kite_sql_serde_macros/src/reference_serialization.rs @@ -127,7 +127,7 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { quote! { impl crate::serdes::ReferenceSerialization for #struct_name { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -141,7 +141,7 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &crate::serdes::ReferenceTables, @@ -208,7 +208,7 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { quote! { impl crate::serdes::ReferenceSerialization for #struct_name { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -222,7 +222,7 @@ pub(crate) fn handle(ast: DeriveInput) -> Result { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &crate::serdes::ReferenceTables, diff --git a/src/binder/aggregate.rs b/src/binder/aggregate.rs index c4226bd4..e17dc898 100644 --- a/src/binder/aggregate.rs +++ b/src/binder/aggregate.rs @@ -16,9 +16,10 @@ 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::MetaArena; use crate::planner::{ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; use crate::{ expression::ScalarExpression, planner::operator::{aggregate::AggregateOperator, sort::SortField}, @@ -28,8 +29,8 @@ struct AggregateCallCollector<'a> { agg_calls: &'a mut Vec, } -impl ExprVisitor> for AggregateCallCollector<'_> { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { +impl ExprVisitor for AggregateCallCollector<'_> { + fn visit(&mut self, expr: ExprRef, arena: &(dyn MetaArena + '_)) -> Result<(), DatabaseError> { match arena.expression(expr) { ScalarExpression::AggCall { .. } => self.agg_calls.push(expr), ScalarExpression::Alias { expr, .. } => self.visit(*expr, arena)?, @@ -40,7 +41,7 @@ impl ExprVisitor> for AggregateCallCollector<'_> { } } -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub fn bind_aggregate( &mut self, children: LogicalPlan, @@ -214,7 +215,7 @@ impl> Binder<'_, '_, T, A> &mut self, select_list: &mut [ExprRef], expr: ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { alias, .. } = arena.expression(expr) { if let Some(i) = select_list.iter().position(|inner_expr| { @@ -275,7 +276,7 @@ impl<'a> HavingOrderByValidator<'a> { } } - fn agg_miss(expr: ExprRef, arena: &PlanArena<'_>) -> DatabaseError { + fn agg_miss(expr: ExprRef, arena: &dyn MetaArena) -> DatabaseError { DatabaseError::AggMiss(format!( "expression '{}' must appear in the GROUP BY clause or be used in an aggregate function", expr.output_name(arena) @@ -283,8 +284,8 @@ impl<'a> HavingOrderByValidator<'a> { } } -impl ExprVisitor> for HavingOrderByValidator<'_> { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { +impl ExprVisitor for HavingOrderByValidator<'_> { + fn visit(&mut self, expr: ExprRef, arena: &(dyn MetaArena + '_)) -> Result<(), DatabaseError> { let contains = |expressions: &[ExprRef]| { expressions .iter() @@ -334,7 +335,7 @@ impl<'a> AggregateOutputBinder<'a> { fn output_ref( &mut self, expr: ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut dyn MetaArena, ) -> Result, DatabaseError> { let output_count = self.agg_calls.len() + self.group_by_exprs.len(); self.agg_calls @@ -370,7 +371,7 @@ impl ExprVisitorMut for AggregateOutputBinder<'_> { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { alias: crate::expression::AliasType::Name(_), @@ -401,7 +402,6 @@ mod tests { use crate::expression::{AliasType, BinaryOperator, ScalarExpression}; use crate::planner::{ExprRef, PlanArena}; use crate::storage::Storage; - use crate::types::value::DataValue; use crate::types::LogicalType; fn test_column(arena: &mut PlanArena, name: &str, ty: LogicalType) -> ColumnRef { @@ -504,7 +504,7 @@ mod tests { let scala_functions = Default::default(); let table_functions = Default::default(); let transaction = tables.storage.transaction()?; - let args: [(&'static str, DataValue); 0] = []; + let args: [(usize, LogicalType); 0] = []; let mut binder = Binder::new( BinderContext::new( &tables.table_cache, diff --git a/src/binder/alter_table.rs b/src/binder/alter_table.rs index b18e27f8..2cb1d303 100644 --- a/src/binder/alter_table.rs +++ b/src/binder/alter_table.rs @@ -23,10 +23,9 @@ use crate::planner::operator::alter_table::drop_column::DropColumnOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_add_column( &mut self, table_name: TableName, diff --git a/src/binder/analyze.rs b/src/binder/analyze.rs index 8a4faf73..2021329d 100644 --- a/src/binder/analyze.rs +++ b/src/binder/analyze.rs @@ -20,9 +20,9 @@ use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_analyze( &mut self, table_name: TableName, diff --git a/src/binder/copy.rs b/src/binder/copy.rs index 0dfff7cb..0e292ac2 100644 --- a/src/binder/copy.rs +++ b/src/binder/copy.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +use crate::types::LogicalType; use std::path::PathBuf; use std::str::FromStr; @@ -64,7 +65,7 @@ impl FromStr for ExtSource { } } -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(super) fn bind_copy_to_file( &mut self, target: ExtSource, diff --git a/src/binder/create_index.rs b/src/binder/create_index.rs index 621f48a0..24702eda 100644 --- a/src/binder/create_index.rs +++ b/src/binder/create_index.rs @@ -21,9 +21,9 @@ use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; use crate::types::index::IndexType; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_create_index_source( &mut self, table_name: TableName, diff --git a/src/binder/create_table.rs b/src/binder/create_table.rs index f4e56a56..9de51118 100644 --- a/src/binder/create_table.rs +++ b/src/binder/create_table.rs @@ -19,10 +19,10 @@ use crate::planner::operator::create_table::CreateTableOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; use std::collections::HashSet; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { // TODO: TableConstraint pub(crate) fn bind_create_table( &mut self, diff --git a/src/binder/create_view.rs b/src/binder/create_view.rs index 4b12e320..42a397aa 100644 --- a/src/binder/create_view.rs +++ b/src/binder/create_view.rs @@ -21,9 +21,9 @@ use crate::planner::operator::create_view::CreateViewOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_create_view( &mut self, view_name: TableName, diff --git a/src/binder/delete.rs b/src/binder/delete.rs index 51b732c2..4eb6617b 100644 --- a/src/binder/delete.rs +++ b/src/binder/delete.rs @@ -19,9 +19,9 @@ use crate::planner::operator::delete::DeleteOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_delete( &mut self, table_name: TableName, diff --git a/src/binder/describe.rs b/src/binder/describe.rs index 440c97f7..3729331d 100644 --- a/src/binder/describe.rs +++ b/src/binder/describe.rs @@ -19,9 +19,9 @@ use crate::planner::operator::describe::DescribeOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_describe( &mut self, table_name: TableName, diff --git a/src/binder/distinct.rs b/src/binder/distinct.rs index 3ce8927a..fb95a97b 100644 --- a/src/binder/distinct.rs +++ b/src/binder/distinct.rs @@ -18,11 +18,12 @@ 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::{ExprRef, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub fn bind_distinct( &mut self, children: LogicalPlan, @@ -83,7 +84,7 @@ impl<'a> DistinctOutputBinder<'a> { Self { select_list } } - fn output_ref(&mut self, expr: ExprRef, arena: &mut PlanArena<'_>) -> Option { + fn output_ref(&mut self, expr: ExprRef, arena: &mut dyn MetaArena) -> Option { self.select_list .iter() .position(|candidate| { @@ -103,7 +104,7 @@ impl ExprVisitorMut for DistinctOutputBinder<'_> { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let ScalarExpression::Alias { alias: crate::expression::AliasType::Name(_), diff --git a/src/binder/drop_index.rs b/src/binder/drop_index.rs index ce138680..149b61ea 100644 --- a/src/binder/drop_index.rs +++ b/src/binder/drop_index.rs @@ -19,9 +19,9 @@ use crate::planner::operator::drop_index::DropIndexOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_drop_index( &mut self, table_name: TableName, diff --git a/src/binder/drop_table.rs b/src/binder/drop_table.rs index 749911f4..047c90fc 100644 --- a/src/binder/drop_table.rs +++ b/src/binder/drop_table.rs @@ -19,9 +19,9 @@ use crate::planner::operator::drop_table::DropTableOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_drop_table( &mut self, table_name: TableName, diff --git a/src/binder/drop_view.rs b/src/binder/drop_view.rs index 7cb85d86..936c6490 100644 --- a/src/binder/drop_view.rs +++ b/src/binder/drop_view.rs @@ -19,9 +19,9 @@ use crate::planner::operator::drop_view::DropViewOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_drop_view( &mut self, view_name: TableName, diff --git a/src/binder/explain.rs b/src/binder/explain.rs index 5b363527..0bd109af 100644 --- a/src/binder/explain.rs +++ b/src/binder/explain.rs @@ -17,9 +17,9 @@ use crate::errors::DatabaseError; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_explain(&mut self, plan: LogicalPlan) -> Result { Ok(LogicalPlan::new( Operator::Explain, diff --git a/src/binder/expr.rs b/src/binder/expr.rs index a0e487ec..85e40a86 100644 --- a/src/binder/expr.rs +++ b/src/binder/expr.rs @@ -38,7 +38,7 @@ macro_rules! try_default { }; } -impl<'a, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, '_, T, A> { +impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A> { fn find_column_in_schema<'schema>( schema_ref: impl IntoIterator, arena: &PlanArena, diff --git a/src/binder/insert.rs b/src/binder/insert.rs index 120a81ea..124561d7 100644 --- a/src/binder/insert.rs +++ b/src/binder/insert.rs @@ -18,21 +18,22 @@ use crate::errors::DatabaseError; use crate::planner::operator::insert::InsertOperator; 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::Transaction; use crate::types::tuple::Schema; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_insert_values( &mut self, table_name: TableName, schema_ref: Schema, - rows: Vec>, + rows: Vec, + row_count: usize, is_overwrite: bool, is_mapping_by_name: bool, ) -> Result { - let values_plan = self.bind_values(rows, schema_ref); + let values_plan = self.bind_values(rows, row_count, schema_ref); Ok(LogicalPlan::new( Operator::Insert(InsertOperator { @@ -62,11 +63,12 @@ impl> Binder<'_, '_, T, A> pub(crate) fn bind_values( &mut self, - rows: Vec>, + rows: Vec, + row_count: usize, schema_ref: Schema, ) -> LogicalPlan { LogicalPlan::new( - Operator::Values(ValuesOperator { rows, schema_ref }), + Operator::Values(ValuesOperator::new(rows, row_count, schema_ref)), Childrens::None, ) } diff --git a/src/binder/mod.rs b/src/binder/mod.rs index edd0ca01..8e61b81b 100644 --- a/src/binder/mod.rs +++ b/src/binder/mod.rs @@ -52,7 +52,9 @@ mod update; mod window; #[cfg(feature = "parser")] -pub use parser::{command_type, prepare, prepare_all, CommandType, Statement}; +pub(crate) use parser::parse_statement; +#[cfg(feature = "parser")] +pub use parser::{command_type, prepare_all, CommandType, Statement}; #[cfg(feature = "orm")] pub use select::{BindPlanFrom, BindPlanSelectList}; #[cfg(feature = "orm")] @@ -69,7 +71,6 @@ use crate::planner::operator::mark_apply::MarkApplyQuantifier; use crate::planner::{ExprRef, LogicalPlan, PlanArena, PlanRef}; use crate::storage::{TableCache, Transaction, ViewCache}; use crate::types::tuple::Schema; -use crate::types::value::DataValue; use crate::types::LogicalType; pub enum InputRefType { @@ -612,7 +613,7 @@ impl<'a, T: Transaction> BinderContext<'a, T> { } } -pub struct Binder<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> { +pub struct Binder<'a, 'parent, T: Transaction, A: AsRef<[(usize, LogicalType)]>> { pub(crate) context: BinderContext<'a, T>, pub(crate) args: &'a A, pub(crate) force_spill: bool, @@ -621,7 +622,7 @@ pub struct Binder<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValu pub(crate) parent: Option<&'parent BinderContext<'a, T>>, } -impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, 'parent, T, A> { +impl<'a, 'parent, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, 'parent, T, A> { pub fn new( context: BinderContext<'a, T>, args: &'a A, diff --git a/src/binder/parser.rs b/src/binder/parser.rs index 804e025f..ae8a760f 100644 --- a/src/binder/parser.rs +++ b/src/binder/parser.rs @@ -24,7 +24,6 @@ use crate::db::{BindSource, DBTransaction, Database, DatabaseIter, TransactionIt use crate::errors::{DatabaseError, SqlErrorSpan}; use crate::expression; 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, TypeCast}; @@ -37,6 +36,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::MetaArena; use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::storage::{Storage, Transaction}; use crate::types::value::{DataValue, Utf8Type}; @@ -54,7 +54,7 @@ pub(super) use sqlparser::ast::{ #[cfg(feature = "copy")] pub(super) use sqlparser::ast::{CopyOption, CopySource, CopyTarget}; use sqlparser::tokenizer::Span; -use std::borrow::{Borrow, Cow}; +use std::borrow::Cow; use std::cmp; use std::slice; @@ -138,7 +138,7 @@ pub fn command_type(stmt: &Statement) -> Result { } /// Parses a single SQL statement into a reusable [`Statement`]. -pub fn prepare>(sql: T) -> Result { +pub(crate) fn parse_statement>(sql: T) -> Result { let mut stmts = prepare_all(sql)?; stmts.pop().ok_or(DatabaseError::EmptyStatement) } @@ -160,24 +160,19 @@ fn statement_mutates_catalog_or_statistics(statement: &Statement) -> Result Database { - /// Executes a prepared [`Statement`] inside a database-owned transaction. - pub fn execute( - &self, - statement: St, - params: A, - ) -> Result, DatabaseError> - where - A: AsRef<[(&'static str, DataValue)]>, - St: Borrow, - { - if statement_mutates_catalog_or_statistics(statement.borrow())? { + /// Bind parameters in a cloned arena and execute an already prepared plan. + pub fn execute<'a>( + &'a self, + prepared: &'a crate::db::PreparedPlan<'_>, + params: impl AsRef<[(usize, DataValue)]>, + ) -> Result, DatabaseError> { + if !std::ptr::eq(prepared.arena.table_arena_cell(), self.state.table_arena()) { return Err(DatabaseError::UnsupportedStmt( - "DDL and ANALYZE require `Database::ddl` or `Database::analyze`".to_string(), + "plan belongs to another database".into(), )); } - BindSource::execute(self, params, |binder, arena| { - binder.bind(statement.borrow(), arena) - }) + let (plan, arena) = prepared.bind_parameters(params.as_ref())?; + BindSource::execute(self, |_, _| Ok((plan, arena))) } pub fn ddl>(&mut self, sql: T) -> Result<(), DatabaseError> { @@ -240,18 +235,20 @@ impl Database { let mut statements = statements.into_iter().peekable(); while let Some(statement) = statements.next() { - let (schema, plan_arena, executor) = - match self + let (schema, plan_arena, executor) = match (|| { + let tx = unsafe { &mut *transaction }; + tx.begin_statement_scope()?; + let (plan, arena) = self .state - .execute(unsafe { &mut *transaction }, &[], |binder, arena| { - binder.bind(&statement, arena) - }) { - Ok(result) => result, - Err(err) => { - unsafe { drop(Box::from_raw(transaction)) }; - return Err(err.with_sql_context(sql)); - } - }; + .build_plan(&[], tx, |binder, arena| binder.bind(&statement, arena))?; + self.state.execute(tx, plan, arena) + })() { + Ok(result) => result, + Err(err) => { + unsafe { drop(Box::from_raw(transaction)) }; + return Err(err.with_sql_context(sql)); + } + }; if statements.peek().is_some() { if let Err(err) = @@ -263,7 +260,7 @@ impl Database { } else { let inner = Box::into_raw(Box::new(TransactionIter::new( schema, - plan_arena, + Box::new(plan_arena) as Box, executor, transaction, ))); @@ -277,27 +274,19 @@ impl Database { } impl<'txn, S: Storage> DBTransaction<'txn, S> { - /// Executes a prepared [`Statement`] inside the current transaction. - pub fn execute<'a, A, St>( + /// Bind parameters in a cloned arena and execute an already prepared plan. + pub fn execute<'a>( &'a mut self, - statement: St, - params: A, - ) -> Result>, DatabaseError> - where - A: AsRef<[(&'static str, DataValue)]>, - St: Borrow, - { - if matches!( - command_type(statement.borrow())?, - CommandType::DDL | CommandType::Analyze - ) { + prepared: &'a crate::db::PreparedPlan<'txn>, + params: impl AsRef<[(usize, DataValue)]>, + ) -> Result>, DatabaseError> { + if !std::ptr::eq(prepared.arena.table_arena_cell(), self.state.table_arena()) { return Err(DatabaseError::UnsupportedStmt( - "`DDL` and `ANALYZE` are not allowed to execute within a transaction".to_string(), + "plan belongs to another database".into(), )); } - BindSource::execute(self, params, |binder, arena| { - binder.bind(statement.borrow(), arena) - }) + let (plan, arena) = prepared.bind_parameters(params.as_ref())?; + BindSource::execute(self, |_, _| Ok((plan, arena))) } /// Runs SQL inside the current transaction and returns the final result iterator. @@ -307,26 +296,38 @@ impl<'txn, S: Storage> DBTransaction<'txn, S> { ) -> Result>, DatabaseError> { let sql = sql.as_ref(); let mut statements = prepare_all(sql).map_err(|err| err.with_sql_context(sql))?; + for statement in &statements { + if statement_mutates_catalog_or_statistics(statement)? { + return Err(DatabaseError::UnsupportedStmt( + "DDL and ANALYZE are not allowed to execute within a transaction".into(), + ) + .with_sql_context(sql)); + } + } let last_statement = statements .pop() .ok_or_else(|| DatabaseError::EmptyStatement.with_sql_context(sql))?; for statement in statements { - self.execute(&statement, &[]) - .map_err(|err| err.with_sql_context(sql))? - .done() - .map_err(|err| err.with_sql_context(sql))?; + BindSource::execute(&mut *self, |state, tx| { + state.build_plan(&[], tx, |binder, arena| binder.bind(&statement, arena)) + }) + .map_err(|err| err.with_sql_context(sql))? + .done() + .map_err(|err| err.with_sql_context(sql))?; } - self.execute(&last_statement, &[]) - .map_err(|err| err.with_sql_context(sql)) + BindSource::execute(self, |state, tx| { + state.build_plan(&[], tx, |binder, arena| binder.bind(&last_statement, arena)) + }) + .map_err(|err| err.with_sql_context(sql)) } } struct BindStatementStart<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut PlanArena<'arena>, @@ -345,7 +346,7 @@ impl ExprVisitorMut for UpdateExprTargetRemapper<'_> { &mut self, column: &mut ColumnRef, position: &mut usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let Some(target_position) = self .target_schema @@ -365,7 +366,7 @@ impl ExprVisitorMut for UpdateExprTargetRemapper<'_> { impl<'s, 'a: 'b, 'b, 'arena, T, A> BindStatementStart<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn statement(self, stmt: &Statement) -> Result { let span = stmt.span(); @@ -1007,7 +1008,7 @@ where values_len, )); } - source.schema().to_vec() + source.schema()[..values_len].to_vec() } else { let mut columns = Vec::with_capacity(idents.len()); for ident in idents { @@ -1026,64 +1027,25 @@ where } columns }; - let mut rows = Vec::with_capacity(expr_rows.len()); - + let mut rows = Vec::with_capacity(expr_rows.len() * values_len); for expr_row in expr_rows { if expr_row.len() != values_len { return Err(DatabaseError::ValuesLenMismatch(expr_row.len(), values_len)); } - let mut row = Vec::with_capacity(expr_row.len()); - for (i, expr) in expr_row.iter().enumerate() { - let span = expr.span(); - let mut expr_ref = self - .binder - .bind_expr(expr, self.arena) - .map(|expr| self.arena.alloc_expression(expr))?; - - 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]); - let ty = column.datatype(); - - value = value.cast(ty)?; - value.check_len(ty)?; - if value.is_null() && !column.nullable() { - return Err(attach_span_if_absent( - DatabaseError::not_null_column(column.name().to_string()), - span, - )); - } - - row.push(value); - } - ScalarExpression::Empty => { - let column = self.arena.column(schema_ref[i]); - let default_value = column + let expression = self.binder.bind_expr(expr, self.arena)?; + let expression = if matches!(expression, ScalarExpression::Empty) { + ScalarExpression::Constant( + self.arena + .column(schema_ref[i]) .default_value(self.arena)? - .ok_or(DatabaseError::DefaultNotExist)?; - if default_value.is_null() && !column.nullable() { - return Err(attach_span_if_absent( - DatabaseError::not_null_column(column.name().to_string()), - span, - )); - } - row.push(default_value); - } - _ => { - return Err(attach_span_if_absent( - DatabaseError::UnsupportedStmt( - "INSERT values must be constants or DEFAULT".to_string(), - ), - span, - )) - } - } + .ok_or(DatabaseError::DefaultNotExist)?, + ) + } else { + expression + }; + rows.push(self.arena.alloc_expression(expression)); } - rows.push(row); } self.binder.context.allow_default = false; @@ -1091,6 +1053,7 @@ where table_name, schema_ref, rows, + expr_rows.len(), is_overwrite, is_mapping_by_name, ) @@ -1261,7 +1224,7 @@ where match (names.next(), exprs.next()) { (Some(name), Some(expression)) => { let expression = std::mem::replace( - self.arena.expression_mut(expression), + &mut *self.arena.expression_mut(expression), ScalarExpression::Empty, ); bind_assignment(self.binder, self.arena, name, expression)? @@ -1371,7 +1334,7 @@ impl BindStatementComplete { impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanStart<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { #[allow(clippy::wrong_self_convention)] pub(crate) fn from_sql( @@ -1404,7 +1367,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanFrom<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn select_list_from_sql( self, @@ -1421,7 +1384,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanSelectList<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn where_sql( self, @@ -1443,7 +1406,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanFiltered<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn aggregate_sql( self, @@ -1500,7 +1463,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> super::select::BindPlanWindowed<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn distinct_sql( self, @@ -1513,7 +1476,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanProjected<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn select_into_sql( self, @@ -1946,7 +1909,7 @@ fn copy_file_format(options: Vec) -> Result> Binder<'a, 'parent, T, A> { +impl<'a, 'parent, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, 'parent, T, A> { fn bind_table_ref_sql( &mut self, from: &TableWithJoins, @@ -2119,17 +2082,27 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< self.bind_binary_op_expr(left_expr, right_expr, op.clone().try_into()?, arena) } Expr::Value(v) => { - let value = if let Value::Placeholder(name) = &v.value { - self.args + let value = if let Value::Placeholder(placeholder) = &v.value { + let id = DataValue::parameter_id(placeholder).ok_or_else(|| { + attach_span_if_absent( + DatabaseError::InvalidValue(format!( + "parameter must use the $ form: {placeholder}" + )), + v, + ) + })?; + let ty = self + .args .as_ref() .iter() - .find_map(|(key, value)| (key == name).then(|| value.clone())) + .find_map(|(candidate, ty)| (candidate == &id).then_some(ty)) .ok_or_else(|| { attach_span_if_absent( - DatabaseError::parameter_not_found(name.to_string()), + DatabaseError::parameter_not_found(format!("${id}")), v, ) - })? + })?; + DataValue::Parameter { id, ty: ty.clone() } } else { (&v.value) .try_into() @@ -2765,40 +2738,24 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< let values_len = expr_rows[0].len(); let mut inferred_types: Vec> = vec![None; values_len]; - let mut rows = Vec::with_capacity(expr_rows.len()); + let mut rows = Vec::with_capacity(expr_rows.len() * values_len); for expr_row in expr_rows { if expr_row.len() != values_len { return Err(DatabaseError::ValuesLenMismatch(expr_row.len(), values_len)); } - 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) - .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(); - - inferred_types[col_index] = match &inferred_types[col_index] { - Some(existing) => { - Some(LogicalType::max_logical_type(existing, &value_type)?.into_owned()) - } - None => Some(value_type), - }; - - row.push(value); - } else { - return Err(DatabaseError::ColumnsEmpty); - } + let expression = self.bind_expr(expr, arena)?; + let value_type = expression.return_type(arena).into_owned(); + inferred_types[col_index] = match &inferred_types[col_index] { + Some(existing) => { + Some(LogicalType::max_logical_type(existing, &value_type)?.into_owned()) + } + None => Some(value_type), + }; + rows.push(arena.alloc_expression(expression)); } - - rows.push(row); } let value_name = arena.temp_table(); @@ -2817,7 +2774,7 @@ impl<'a, 'parent, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder< }) .collect::>()?; - Ok(self.bind_values(rows, column_refs)) + Ok(self.bind_values(rows, expr_rows.len(), column_refs)) } fn bind_top_level_orderby( @@ -3266,36 +3223,57 @@ mod tests { } } + #[test] + fn prepare_rejects_non_positional_parameters() -> Result<(), DatabaseError> { + let db = DataBaseBuilder::path(".").build_in_memory()?; + for placeholder in ["$0", "$01", "$name"] { + assert!(matches!( + db.prepare( + &format!("select {placeholder}"), + &[(1, LogicalType::Integer)] + ), + Err(DatabaseError::InvalidValue(_)) + )); + } + Ok(()) + } + #[test] fn test_prepare_and_command_type_classification() -> Result<(), DatabaseError> { assert!(matches!( prepare_all(""), Err(DatabaseError::EmptyStatement) )); - assert_eq!(command_type(&prepare("select 1")?)?, CommandType::DQL); assert_eq!( - command_type(&prepare("create table t (id int primary key)")?)?, + command_type(&parse_statement("select 1")?)?, + CommandType::DQL + ); + assert_eq!( + command_type(&parse_statement("create table t (id int primary key)")?)?, CommandType::DDL ); assert_eq!( - command_type(&prepare("analyze table t")?)?, + command_type(&parse_statement("analyze table t")?)?, CommandType::Analyze ); assert_eq!( - command_type(&prepare("insert into t values (1)")?)?, + command_type(&parse_statement("insert into t values (1)")?)?, CommandType::DML ); assert_eq!( - command_type(&prepare("update t set id = 1")?)?, + command_type(&parse_statement("update t set id = 1")?)?, CommandType::DML ); - assert_eq!(command_type(&prepare("delete from t")?)?, CommandType::DML); assert_eq!( - command_type(&prepare("truncate table t")?)?, + command_type(&parse_statement("delete from t")?)?, + CommandType::DML + ); + assert_eq!( + command_type(&parse_statement("truncate table t")?)?, CommandType::DML ); - let err = command_type(&prepare("start transaction")?).unwrap_err(); + let err = command_type(&parse_statement("start transaction")?).unwrap_err(); assert_unsupported(err, "START TRANSACTION"); Ok(()) @@ -3304,11 +3282,10 @@ mod tests { #[test] fn test_database_entrypoints_reject_catalog_mutation() -> Result<(), DatabaseError> { let mut database = DataBaseBuilder::path(".").build_in_memory()?; - let params = &[] as &[(&'static str, DataValue)]; + let params = &[] as &[(usize, LogicalType)]; - let ddl = prepare("create table t (id int primary key)")?; assert_unsupported( - expect_err(database.execute(&ddl, params)), + expect_err(database.prepare("create table t (id int primary key)", params)), "DDL and ANALYZE", ); assert_unsupported( @@ -3317,11 +3294,10 @@ mod tests { ); assert_unsupported(expect_err(database.ddl("select 1")), "`Database::ddl`"); - let analyze = prepare("analyze table t")?; - let mut transaction = database.new_transaction()?; + let transaction = database.new_transaction()?; assert_unsupported( - expect_err(transaction.execute(&analyze, params)), - "not allowed to execute within a transaction", + expect_err(transaction.prepare("analyze table t", params)), + "DDL and ANALYZE", ); transaction.commit()?; @@ -3375,7 +3351,7 @@ mod tests { "only a single ALTER TABLE operation", ); - let mut stmt = prepare("alter table t1 drop column c1")?; + let mut stmt = parse_statement("alter table t1 drop column c1")?; let Statement::AlterTable(alter) = &mut stmt else { unreachable!("expected alter table statement") }; diff --git a/src/binder/select.rs b/src/binder/select.rs index da6875e1..6258c0fb 100644 --- a/src/binder/select.rs +++ b/src/binder/select.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +use crate::planner::MetaArena; use crate::{ expression::ScalarExpression, planner::{ @@ -22,7 +23,6 @@ use crate::{ }, operator::{join::JoinType, table_scan::TableScanOperator}, }, - types::value::DataValue, }; use std::{borrow::Cow, collections::HashSet}; @@ -55,7 +55,7 @@ impl ExprVisitorMut for RightSidePositionGlobalizer<'_> { &mut self, column: &mut ColumnRef, position: &mut usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if self .right_schema @@ -84,7 +84,7 @@ impl ExprVisitorMut for SplitScopePositionRebinder<'_> { &mut self, column: &mut ColumnRef, position: &mut usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let Some(left_position) = self .left_schema @@ -113,7 +113,7 @@ impl ExprVisitorMut for MarkerPositionGlobalizer<'_> { &mut self, column: &mut ColumnRef, position: &mut usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if arena.same_column(*column, *self.output_column) { *position = self.left_len; @@ -131,7 +131,7 @@ impl<'a> ProjectionOutputBinder<'a> { Self { project_exprs } } - fn output_ref(&mut self, expr: ExprRef, arena: &mut PlanArena<'_>) -> Option { + fn output_ref(&mut self, expr: ExprRef, arena: &mut dyn MetaArena) -> Option { self.project_exprs .iter() .position(|candidate| { @@ -151,7 +151,7 @@ impl ExprVisitorMut for ProjectionOutputBinder<'_> { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let Some(output_ref) = self.output_ref(*expr, arena) { *expr = arena.alloc_expression(output_ref); @@ -164,7 +164,7 @@ impl ExprVisitorMut for ProjectionOutputBinder<'_> { pub(crate) struct BindPlanStart<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) binder: &'s mut Binder<'a, 'b, T, A>, pub(crate) arena: &'s mut crate::planner::PlanArena<'arena>, @@ -173,7 +173,7 @@ where pub struct BindPlanFrom<'s, 'a, 'b, 'arena, T, A, M = ()> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) binder: &'s mut Binder<'a, 'b, T, A>, pub(crate) arena: &'s mut crate::planner::PlanArena<'arena>, @@ -184,7 +184,7 @@ where pub struct BindPlanSelectList<'s, 'a, 'b, 'arena, T, A, M = ()> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) binder: &'s mut Binder<'a, 'b, T, A>, pub(crate) arena: &'s mut crate::planner::PlanArena<'arena>, @@ -196,7 +196,7 @@ where pub(crate) struct BindPlanFiltered<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(super) binder: &'s mut Binder<'a, 'b, T, A>, pub(super) arena: &'s mut crate::planner::PlanArena<'arena>, @@ -207,7 +207,7 @@ where pub(crate) struct BindPlanAggregated<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -220,7 +220,7 @@ where pub(crate) struct BindPlanHaving<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -232,7 +232,7 @@ where pub(crate) struct BindPlanWindowed<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -244,7 +244,7 @@ where pub(crate) struct BindPlanDistinct<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -256,7 +256,7 @@ where pub(crate) struct BindPlanSorted<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'s mut Binder<'a, 'b, T, A>, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -267,7 +267,7 @@ where pub(crate) struct BindPlanProjected<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { plan: LogicalPlan, _marker: std::marker::PhantomData<(&'s (), &'a (), &'b (), &'arena (), T, A)>, @@ -292,7 +292,7 @@ pub(crate) enum JoinConstraintInput { impl<'s, 'a: 'b, 'b, 'arena, T, A, M> BindPlanFrom<'s, 'a, 'b, 'arena, T, A, M> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { #[cfg(feature = "orm")] pub(crate) fn typed(self) -> BindPlanFrom<'s, 'a, 'b, 'arena, T, A, N> { @@ -344,7 +344,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A, M> BindPlanSelectList<'s, 'a, 'b, 'arena, T, A, M> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { #[cfg(feature = "orm")] pub(crate) fn set_select_list(mut self, select_list: Vec) -> Self { @@ -453,7 +453,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanStart<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { #[allow(clippy::wrong_self_convention)] pub(crate) fn from_plan( @@ -472,7 +472,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A, M> BindPlanSelectList<'s, 'a, 'b, 'arena, T, A, M> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn filter_expr( mut self, @@ -496,7 +496,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanFiltered<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn aggregate( mut self, @@ -571,7 +571,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanAggregated<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn having( mut self, @@ -593,7 +593,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanHaving<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn window( mut self, @@ -618,7 +618,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanWindowed<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn distinct( mut self, @@ -651,7 +651,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanDistinct<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn order_by( mut self, @@ -672,7 +672,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanSorted<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn project( mut self, @@ -693,7 +693,7 @@ where impl<'s, 'a: 'b, 'b, 'arena, T, A> BindPlanProjected<'s, 'a, 'b, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub(crate) fn insert_into( mut self, @@ -720,7 +720,7 @@ impl BindPlanComplete { } } -impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'a, 'b, T, A> { +impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, 'b, T, A> { pub(crate) fn build_plan<'s, 'arena>( &'s mut self, arena: &'s mut crate::planner::PlanArena<'arena>, @@ -841,7 +841,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' &mut self, column: &mut ColumnRef, position: &mut usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let Some(output) = self.appended_outputs.iter().find(|output| { *position == output.child_position && arena.same_column(*column, output.column) @@ -1973,7 +1973,7 @@ impl<'a: 'b, 'b, T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<' for expr in select_items { let mut expression = - std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + std::mem::replace(&mut *arena.expression_mut(*expr), ScalarExpression::Empty); if let ScalarExpression::ColumnRef { column, .. } = &mut expression { let _ = table_force_nullable .iter() diff --git a/src/binder/show_table.rs b/src/binder/show_table.rs index 347fe2b9..c6a0891c 100644 --- a/src/binder/show_table.rs +++ b/src/binder/show_table.rs @@ -17,9 +17,9 @@ use crate::errors::DatabaseError; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_show_tables(&mut self) -> Result { Ok(LogicalPlan::new(Operator::ShowTable, Childrens::None)) } diff --git a/src/binder/show_view.rs b/src/binder/show_view.rs index 63a81ed2..dc90320e 100644 --- a/src/binder/show_view.rs +++ b/src/binder/show_view.rs @@ -17,9 +17,9 @@ use crate::errors::DatabaseError; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_show_views(&mut self) -> Result { Ok(LogicalPlan::new(Operator::ShowView, Childrens::None)) } diff --git a/src/binder/truncate.rs b/src/binder/truncate.rs index cd81370f..f3e1611d 100644 --- a/src/binder/truncate.rs +++ b/src/binder/truncate.rs @@ -19,9 +19,9 @@ use crate::planner::operator::truncate::TruncateOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_truncate( &mut self, table_name: TableName, diff --git a/src/binder/update.rs b/src/binder/update.rs index 9e1270d9..bab3520b 100644 --- a/src/binder/update.rs +++ b/src/binder/update.rs @@ -19,9 +19,9 @@ use crate::planner::operator::update::UpdateOperator; use crate::planner::operator::Operator; use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::Transaction; -use crate::types::value::DataValue; +use crate::types::LogicalType; -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_update( &mut self, table_name: TableName, diff --git a/src/binder/window.rs b/src/binder/window.rs index fd33e0b6..0cb195b8 100644 --- a/src/binder/window.rs +++ b/src/binder/window.rs @@ -22,9 +22,9 @@ 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::MetaArena; use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::storage::Transaction; -use crate::types::value::DataValue; use crate::types::LogicalType; struct WindowCollector { @@ -35,7 +35,7 @@ impl ExprVisitorMut for WindowCollector { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let ScalarExpression::WindowCall(window) = arena.expression(*expr) else { return walk_mut_expr(self, expr, arena); @@ -51,7 +51,7 @@ impl ExprVisitorMut for WindowCollector { let output_name = expr.output_name(arena); let ScalarExpression::WindowCall(window) = - std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty) + std::mem::replace(&mut *arena.expression_mut(*expr), ScalarExpression::Empty) else { unreachable!() }; @@ -76,7 +76,7 @@ impl ExprVisitorMut for WindowOutputBinder<'_> { &mut self, column: &mut ColumnRef, position: &mut usize, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let Some(output_position) = self .groups @@ -97,7 +97,7 @@ struct WindowGroup { output_columns: Vec, } -impl> Binder<'_, '_, T, A> { +impl> Binder<'_, '_, T, A> { pub(crate) fn bind_window_function( &mut self, kind: WindowFunctionKind, diff --git a/src/catalog/column.rs b/src/catalog/column.rs index 67741ca3..04b20526 100644 --- a/src/catalog/column.rs +++ b/src/catalog/column.rs @@ -14,12 +14,13 @@ use crate::catalog::TableName; use crate::errors::DatabaseError; -use crate::planner::{ExprRef, PlanArena}; -use crate::types::tuple::Tuple; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::value::DataValue; use crate::types::CharLengthUnits; use crate::types::{ColumnId, LogicalType}; use kite_sql_serde_macros::ReferenceSerialization; +use std::borrow::Cow; use std::fmt; use std::hash::Hash; @@ -174,12 +175,17 @@ impl ColumnCatalog { pub(crate) fn default_value( &self, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { self.desc .default .as_ref() - .map(|expr| arena.expression(*expr).eval::<&Tuple>(arena, None)) + .map(|expr| { + arena + .expression(*expr) + .eval(arena, None) + .map(Cow::into_owned) + }) .transpose() } diff --git a/src/catalog/table.rs b/src/catalog/table.rs index 9f8eb0f1..d77f4512 100644 --- a/src/catalog/table.rs +++ b/src/catalog/table.rs @@ -62,7 +62,7 @@ impl TableCatalog { pub(crate) fn get_unique_index( &self, col_id: &ColumnId, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Option { self.indexes.iter().copied().find(|meta| { let meta = arena.index(*meta); @@ -121,14 +121,17 @@ impl TableCatalog { pub(crate) fn dml_snapshot( &self, - arena: &mut PlanArena, + arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { let index_metas = self .indexes() .map(|index_meta| { Ok(( *index_meta, - arena.index(*index_meta).column_exprs(self, arena)?, + arena + .index(*index_meta) + .column_exprs(self) + .collect::, _>>()?, )) }) .collect::, DatabaseError>>()?; @@ -145,7 +148,7 @@ impl TableCatalog { pub(crate) fn add_column( &mut self, mut col: ColumnCatalog, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { if self.column_idxs.contains_key(col.name()) { return Err(DatabaseError::DuplicateColumn(col.name().to_string())); @@ -176,7 +179,7 @@ impl TableCatalog { name: String, column_ids: Vec, ty: IndexType, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { for index in self.indexes.iter() { if arena.index(*index).name == name { @@ -224,7 +227,7 @@ impl TableCatalog { pub fn new( name: TableName, columns: Vec, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { if columns.is_empty() { return Err(DatabaseError::ColumnsEmpty); @@ -254,7 +257,7 @@ impl TableCatalog { fn build_primary_key_type( primary_keys: &[(usize, ColumnRef)], - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> LogicalType { if primary_keys.len() == 1 { arena.column(primary_keys[0].1).datatype().clone() @@ -272,7 +275,7 @@ impl TableCatalog { name: TableName, column_catalogs: I, indexes: I2, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result where I: Iterator, @@ -332,7 +335,7 @@ impl TableCatalog { fn build_primary_keys( columns: &[ColumnRef], - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> (Vec<(usize, ColumnRef)>, Vec) { let mut primary_keys = Vec::new(); let mut primary_key_indices = Vec::new(); diff --git a/src/catalog/view.rs b/src/catalog/view.rs index 503c9b7b..ae3834d8 100644 --- a/src/catalog/view.rs +++ b/src/catalog/view.rs @@ -33,7 +33,7 @@ impl View { f: &mut F, ) -> Result<(), crate::errors::DatabaseError> where - A: MetaArena, + A: MetaArena + ?Sized, F: FnMut(&ColumnRef) + ?Sized, { for column in &self.schema { diff --git a/src/db.rs b/src/db.rs index f44a2229..2fef0c0c 100644 --- a/src/db.rs +++ b/src/db.rs @@ -13,7 +13,9 @@ // limitations under the License. #[cfg(feature = "parser")] -pub use crate::binder::{prepare, prepare_all, Statement}; +mod prepared; +#[cfg(feature = "parser")] +pub use crate::binder::{prepare_all, Statement}; use crate::binder::{Binder, BinderContext}; use crate::catalog::TableName; use crate::errors::DatabaseError; @@ -40,7 +42,7 @@ use crate::optimizer::rule::normalization::NormalizationRuleImpl; #[cfg(feature = "orm")] use crate::orm::FromQueryRow; use crate::planner::operator::Operator; -use crate::planner::{LogicalPlan, PlanArena, TableArenaCell}; +use crate::planner::{LogicalPlan, MetaArena, PlanArena, TableArenaCell}; #[cfg(all(not(target_arch = "wasm32"), feature = "lmdb"))] use crate::storage::lmdb::{LmdbConfig, LmdbStorage}; use crate::storage::memory::MemoryStorage; @@ -53,6 +55,9 @@ use crate::storage::{ }; use crate::types::tuple::{Schema, SchemaView, Tuple}; use crate::types::value::DataValue; +use crate::types::LogicalType; +#[cfg(feature = "parser")] +pub use prepared::PreparedPlan; use std::collections::{HashMap, HashSet}; use std::marker::PhantomData; use std::mem; @@ -70,22 +75,23 @@ pub enum CatalogKind { TableFunction(Arc), } -pub(crate) trait BindSource { +pub(crate) trait BindSource<'a> { type Iter: ResultIter; type Transaction: Transaction; - fn execute(self, params: A, build: F) -> Result + type Storage: Storage; + + fn execute(self, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, - F: for<'bind> FnOnce( - &mut Binder<'bind, '_, Self::Transaction, A>, - &mut PlanArena<'_>, - ) -> Result; + F: FnOnce( + &'a State, + &Self::Transaction, + ) -> Result<(LogicalPlan, A), DatabaseError>; #[cfg(feature = "orm")] fn explain(self, params: A, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, F: for<'bind> FnOnce( &mut Binder<'bind, '_, Self::Transaction, A>, &mut PlanArena<'_>, @@ -284,7 +290,7 @@ impl DataBaseBuilder { table_cache, view_cache, table_arena, - optimizer_pipeline: default_optimizer_pipeline(), + optimizer_pipeline: optimizer_pipeline(), histogram_buckets, _p: Default::default(), }; @@ -308,7 +314,7 @@ impl DataBaseBuilder { } } -fn default_optimizer_pipeline() -> HepOptimizerPipeline { +fn optimizer_pipeline() -> HepOptimizerPipeline { HepOptimizerPipeline::builder() .before_batch( "Column Pruning".to_string(), @@ -331,7 +337,6 @@ fn default_optimizer_pipeline() -> HepOptimizerPipeline { vec![ NormalizationRuleImpl::PushPredicateThroughJoin, NormalizationRuleImpl::PushJoinPredicateIntoScan, - NormalizationRuleImpl::PushPredicateIntoScan, ], ) .before_batch( @@ -359,6 +364,11 @@ fn default_optimizer_pipeline() -> HepOptimizerPipeline { NormalizationRuleImpl::CombineFilter, ], ) + .before_batch( + "Predicate Into Scan".to_string(), + HepBatchStrategy::fix_point_topdown(10), + vec![NormalizationRuleImpl::PushPredicateIntoScan], + ) .after_batch( "Parameterize Mark Apply".to_string(), HepBatchStrategy::once_topdown(), @@ -505,21 +515,20 @@ impl State { Ok(()) } - pub(crate) fn build_plan<'a, 'txn, A: AsRef<[(&'static str, DataValue)]>, F>( + pub(crate) fn build_plan<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>, F>( &'a self, params: A, - transaction: &::TransactionType<'txn>, + transaction: &T, build: F, ) -> Result<(LogicalPlan, PlanArena<'a>), DatabaseError> where - S: 'txn, F: for<'bind> FnOnce( - &mut Binder<'bind, '_, ::TransactionType<'txn>, A>, + &mut Binder<'bind, '_, T, A>, &mut PlanArena<'a>, ) -> Result, { - let mut plan_arena = PlanArena::new(self.table_arena()); - let mut binder: Binder<'_, '_, ::TransactionType<'txn>, A> = Binder::new( + let mut arena = PlanArena::new(self.table_arena()); + let mut binder = Binder::new( BinderContext::new( self.table_cache(), self.view_cache(), @@ -530,63 +539,44 @@ impl State { ¶ms, None, ); - let source_plan = build(&mut binder, &mut plan_arena)?; + let source_plan = build(&mut binder, &mut arena)?; drop(binder); - let mut best_plan = self.optimizer_pipeline.instantiate(source_plan).find_best( + let mut plan = self.optimizer_pipeline.instantiate(source_plan).find_best( Some(&StatisticMetaLoader::new(self.meta_cache())), - &mut plan_arena, + &mut arena, )?; - if let Operator::Analyze(op) = &mut best_plan.operator { + if let Operator::Analyze(op) = &mut plan.operator { if op.histogram_buckets.is_none() { op.histogram_buckets = self.histogram_buckets; } } - Ok((best_plan, plan_arena)) + Ok((plan, arena)) } - pub(crate) fn execute<'a, 'txn, A, F>( + pub(crate) fn execute<'a, 'txn, A: MetaArena + 'a>( &'a self, transaction: &'a mut S::TransactionType<'txn>, - params: A, - build: F, - ) -> Result< - ( - Schema, - PlanArena<'a>, - Executor<'a, S::TransactionType<'txn>>, - ), - DatabaseError, - > + mut plan: LogicalPlan, + mut plan_arena: A, + ) -> Result<(Schema, A, Executor<'a, S::TransactionType<'txn>>), DatabaseError> where S: 'txn, - A: AsRef<[(&'static str, DataValue)]>, - F: for<'bind> FnOnce( - &mut Binder<'bind, '_, S::TransactionType<'txn>, A>, - &mut PlanArena<'a>, - ) -> Result, { - transaction.begin_statement_scope()?; - match (|| { - let (mut plan, mut plan_arena) = self.build_plan(params, transaction, build)?; - let schema = plan.take_schema(&mut plan_arena); - let mut arena = ExecArena::new(); - let read_context = ExecutionContext::new( - &self.table_cache, - &self.view_cache, - &self.meta_cache, - &self.scala_functions, - &self.table_functions, - ); - let root = build_write(&mut arena, &mut plan_arena, plan, read_context, transaction); - let executor = Executor::new(arena, root); + let schema = plan.take_schema(&mut plan_arena); + let mut arena = ExecArena::new(); + let read_context = ExecutionContext::new( + &self.table_cache, + &self.view_cache, + &self.meta_cache, + &self.scala_functions, + &self.table_functions, + ); + let root = build_write(&mut arena, &mut plan_arena, plan, read_context, transaction); + let executor = Executor::new(arena, root); - Ok((schema, plan_arena, executor)) - })() { - Ok(result) => Ok(result), - Err(err) => Err(err), - } + Ok((schema, plan_arena, executor)) } pub(crate) fn execute_mut<'a, 'txn, A, F>( @@ -604,7 +594,7 @@ impl State { > where S: 'txn, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, F: for<'bind> FnOnce( &mut Binder<'bind, '_, S::TransactionType<'txn>, A>, &mut PlanArena<'a>, @@ -734,7 +724,7 @@ impl Database { build: F, ) -> Result<(), DatabaseError> where - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, F: for<'a, 'txn, 'bind> FnOnce( &mut Binder<'bind, '_, S::TransactionType<'txn>, A>, &mut PlanArena<'a>, @@ -875,36 +865,36 @@ impl Database { } } -impl<'a, S: Storage> BindSource for &'a Database { +impl<'a, S: Storage> BindSource<'a> for &'a Database { type Iter = DatabaseIter<'a, S>; type Transaction = S::TransactionType<'a>; - fn execute(self, params: A, build: F) -> Result + type Storage = S; + + fn execute(self, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, - F: for<'bind> FnOnce( - &mut Binder<'bind, '_, Self::Transaction, A>, - &mut PlanArena<'_>, - ) -> Result, + F: FnOnce(&'a State, &Self::Transaction) -> Result<(LogicalPlan, A), DatabaseError>, { let transaction = Box::into_raw(Box::new( self.storage .transaction_with_isolation(self.transaction_isolation)?, )); - let (schema, plan_arena, executor) = - match self - .state - .execute(unsafe { &mut *transaction }, params, build) - { - Ok(result) => result, - Err(err) => { - unsafe { drop(Box::from_raw(transaction)) }; - return Err(err); - } - }; + let result = (|| { + let transaction = unsafe { &mut *transaction }; + transaction.begin_statement_scope()?; + let (plan, arena) = build(&self.state, transaction)?; + self.state.execute(transaction, plan, arena) + })(); + let (schema, arena, executor) = match result { + Ok(result) => result, + Err(error) => { + unsafe { drop(Box::from_raw(transaction)) }; + return Err(error); + } + }; let inner = Box::into_raw(Box::new(TransactionIter::new( schema, - plan_arena, + Box::new(arena) as Box, executor, transaction, ))); @@ -914,7 +904,7 @@ impl<'a, S: Storage> BindSource for &'a Database { #[cfg(feature = "orm")] fn explain(self, params: A, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, F: for<'bind> FnOnce( &mut Binder<'bind, '_, Self::Transaction, A>, &mut PlanArena<'_>, @@ -1093,25 +1083,25 @@ impl<'txn, S: Storage> DBTransaction<'txn, S> { } } -impl<'a, 'txn, S: Storage> BindSource for &'a mut DBTransaction<'txn, S> { +impl<'a, 'txn, S: Storage> BindSource<'a> for &'a mut DBTransaction<'txn, S> { type Iter = TransactionIter<'a, S::TransactionType<'txn>>; type Transaction = S::TransactionType<'txn>; - fn execute(self, params: A, build: F) -> Result + type Storage = S; + + fn execute(self, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, - F: for<'bind> FnOnce( - &mut Binder<'bind, '_, Self::Transaction, A>, - &mut PlanArena<'_>, - ) -> Result, + F: FnOnce(&'a State, &Self::Transaction) -> Result<(LogicalPlan, A), DatabaseError>, { + self.inner.begin_statement_scope()?; + let (plan, arena) = build(self.state, &self.inner)?; let transaction = std::ptr::from_mut(&mut self.inner); - let (schema, plan_arena, executor) = + let (schema, arena, executor) = self.state - .execute(unsafe { &mut *transaction }, params, build)?; + .execute(unsafe { &mut *transaction }, plan, arena)?; Ok(TransactionIter::new( schema, - plan_arena, + Box::new(arena) as Box, executor, transaction, )) @@ -1120,7 +1110,7 @@ impl<'a, 'txn, S: Storage> BindSource for &'a mut DBTransaction<'txn, S> { #[cfg(feature = "orm")] fn explain(self, params: A, build: F) -> Result where - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, F: for<'bind> FnOnce( &mut Binder<'bind, '_, Self::Transaction, A>, &mut PlanArena<'_>, @@ -1133,19 +1123,19 @@ impl<'a, 'txn, S: Storage> BindSource for &'a mut DBTransaction<'txn, S> { } /// Raw result iterator returned by transaction execution APIs. -pub struct TransactionIter<'a, T: Transaction + 'a> { +pub struct TransactionIter<'a, T: Transaction + 'a, A: MetaArena + 'a = Box> { executor: Option>, - plan_arena: Option>, + plan_arena: Option, schema: Schema, transaction: *mut T, statement_scope_active: bool, ddl_apply: Vec, } -impl<'a, T: Transaction + 'a> TransactionIter<'a, T> { +impl<'a, T: Transaction + 'a, A: MetaArena + 'a> TransactionIter<'a, T, A> { pub(crate) fn new( schema: Schema, - plan_arena: PlanArena<'a>, + plan_arena: A, executor: Executor<'a, T>, transaction: *mut T, ) -> Self { @@ -1216,7 +1206,9 @@ impl<'a, T: Transaction + 'a> TransactionIter<'a, T> { while self.next_tuple(|_, _| ())?.is_some() {} Ok(()) } +} +impl<'a, T: Transaction + 'a> TransactionIter<'a, T, PlanArena<'a>> { fn done_with_ddl_apply(mut self) -> Result<(PlanArena<'a>, Vec), DatabaseError> { while self.next_tuple(|_, _| ())?.is_some() {} Ok(( @@ -1228,13 +1220,13 @@ impl<'a, T: Transaction + 'a> TransactionIter<'a, T> { } } -impl Drop for TransactionIter<'_, T> { +impl Drop for TransactionIter<'_, T, A> { fn drop(&mut self) { let _ = self.finish_statement_scope(); } } -impl ResultIter for TransactionIter<'_, T> { +impl ResultIter for TransactionIter<'_, T, A> { fn schema(&self, f: impl FnOnce(&SchemaView<'_, '_>) -> R) -> R { TransactionIter::schema(self, f) } @@ -1452,7 +1444,7 @@ pub(crate) mod test { kite_sql.ddl("CREATE TABLE onecolumn (id INT PRIMARY KEY, x INT NULL)")?; kite_sql.ddl("CREATE TABLE empty (e_id INT PRIMARY KEY, x INT)")?; - let stmt = crate::db::prepare( + let stmt = crate::binder::parse_statement( "SELECT * FROM onecolumn AS a(aid, x) JOIN empty AS b(bid, y) ON a.x = b.y", )?; let transaction = kite_sql.storage.transaction()?; @@ -1535,7 +1527,7 @@ pub(crate) mod test { kite_sql.ddl("CREATE TABLE onecolumn (id INT PRIMARY KEY, x INT NULL)")?; kite_sql.ddl("CREATE TABLE twocolumn (t_id INT PRIMARY KEY, x INT NULL, y INT NULL)")?; - let stmt = crate::db::prepare( + let stmt = crate::binder::parse_statement( "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()?; @@ -1645,7 +1637,7 @@ pub(crate) mod test { )? .done()?; - let stmt = crate::db::prepare( + let stmt = crate::binder::parse_statement( "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()?; @@ -1713,18 +1705,24 @@ pub(crate) mod test { #[test] fn test_prepare_statment() -> 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()?; + let mut kite_sql = DataBaseBuilder::path(temp_dir.path()) + .histogram_buckets(2) + .build_rocksdb()?; kite_sql.ddl("create table t1 (a int primary key, b int)")?; kite_sql.run("insert into t1 values(0, 0)")?.done()?; kite_sql.run("insert into t1 values(1, 1)")?.done()?; kite_sql.run("insert into t1 values(2, 2)")?.done()?; + kite_sql.analyze("t1")?; // Filter { - let statement = crate::db::prepare("explain select * from t1 where b > $1")?; + let statement = kite_sql.prepare( + "explain select * from t1 where b > $1", + &[(1, LogicalType::Integer)], + )?; - let mut iter = kite_sql.execute(statement, &[("$1", DataValue::Int32(0))])?; + let mut iter = kite_sql.execute(&statement, &[(1, DataValue::Int32(0))])?; let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); @@ -1735,43 +1733,48 @@ pub(crate) mod test { } // Aggregate { - let statement = crate::db::prepare( - "explain select a + $1, max(b + $2) from t1 where b > $3 group by a + $4", + let statement = kite_sql.prepare( + "explain select a + $1, max(b + $2) from t1 where b > $3 group by a + $1", + &[ + (1, LogicalType::Integer), + (2, LogicalType::Integer), + (3, LogicalType::Integer), + ], )?; let mut iter = kite_sql.execute( - statement, + &statement, &[ - ("$1", DataValue::Int32(0)), - ("$2", DataValue::Int32(0)), - ("$3", DataValue::Int32(1)), - ("$4", DataValue::Int32(0)), + (1, DataValue::Int32(0)), + (2, DataValue::Int32(0)), + (3, DataValue::Int32(1)), + (4, DataValue::Int32(0)), ], )?; let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); 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)]" + "Projection [(t1.a + $1), Max((t1.b + $2))] [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")?; + let statement = kite_sql.prepare("explain select *, $1 from (select * from t1 where b > $2) left join (select * from t1 where a > $3) on a > $4", &[(1, LogicalType::Integer), (2, LogicalType::Integer), (3, LogicalType::Integer), (4, LogicalType::Integer)])?; let mut iter = kite_sql.execute( - statement, + &statement, &[ - ("$1", DataValue::Int32(9)), - ("$2", DataValue::Int32(0)), - ("$3", DataValue::Int32(1)), - ("$4", DataValue::Int32(0)), + (1, DataValue::Int32(9)), + (2, DataValue::Int32(0)), + (3, DataValue::Int32(1)), + (4, DataValue::Int32(0)), ], )?; let row = next_tuple_owned(&mut iter)?.unwrap(); let plan = row.values[0].utf8().unwrap(); 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)]" + "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)] TableScan t1 -> [t1.a, t1.b] [IndexScan By pk_index => (1, +inf) => (Sort Option: OrderBy: (t1.a Asc Nulls Last) ignore_prefix_len: 0)]" ); } diff --git a/src/db/prepared.rs b/src/db/prepared.rs new file mode 100644 index 00000000..ca8fdf27 --- /dev/null +++ b/src/db/prepared.rs @@ -0,0 +1,563 @@ +use super::*; +use crate::expression::range_detacher::{IndexRangeColumn, RangeDetacher}; +use crate::planner::operator::table_scan::TableScanOperator; +use crate::planner::operator::visitor_mut::OperatorVisitorMut; +use crate::planner::operator::{PhysicalOption, PlanImpl, SortOption}; +use crate::planner::{ExprRef, MetaArena, ParamArena}; +use crate::types::index::IndexLookup; +use crate::types::LogicalType; + +struct ParameterBinder<'a> { + params: &'a [(usize, DataValue)], +} + +impl ParameterBinder<'_> { + fn bind_index_infos( + &self, + index_infos: &mut [crate::types::index::IndexInfo], + ) -> Result<(), DatabaseError> { + for info in index_infos { + if let Some(lookup) = &mut info.lookup { + lookup.bind_parameters(self.params)?; + } + } + Ok(()) + } +} + +impl<'plan> OperatorVisitorMut<'plan> for ParameterBinder<'_> { + fn visit_table_scan( + &mut self, + TableScanOperator { index_infos, .. }: &'plan mut TableScanOperator, + ) -> Result<(), DatabaseError> { + self.bind_index_infos(index_infos) + } + + fn visit_physical_option( + &mut self, + physical_option: &'plan mut PhysicalOption, + ) -> Result<(), DatabaseError> { + if let PlanImpl::IndexScan(info) = &mut physical_option.plan { + self.bind_index_infos(std::slice::from_mut(info))?; + } + Ok(()) + } +} + +/// Extend the selected static index range using its bound residual predicate. +struct SpecializeIndexRange<'a, A: MetaArena + ?Sized> { + arena: &'a mut A, +} + +impl<'plan, A: MetaArena + ?Sized> OperatorVisitorMut<'plan> for SpecializeIndexRange<'_, A> { + fn visit_operator( + &mut self, + operator: &'plan mut Operator, + physical_option: Option<&'plan mut PhysicalOption>, + ) -> Result<(), DatabaseError> { + let Operator::TableScan(scan) = operator else { + return Ok(()); + }; + let Some(option) = physical_option else { + return Ok(()); + }; + let PlanImpl::IndexScan(index) = &mut option.plan else { + return Ok(()); + }; + let (Some(params_predicate), Some(IndexLookup::Static(original))) = + (index.residual_predicate, &index.lookup) + else { + return Ok(()); + }; + let SortOption::OrderBy { + ignore_prefix_len, .. + } = &index.sort_option + else { + return Ok(()); + }; + let Some(range) = RangeDetacher::::specialize_range( + index.meta, + original, + params_predicate, + *ignore_prefix_len, + self.arena, + )? + else { + return Ok(()); + }; + index.lookup = Some(IndexLookup::Static(range)); + for candidate in &mut scan.index_infos { + if candidate.meta == index.meta { + candidate.lookup = index.lookup.clone(); + } + } + Ok(()) + } +} + +/// A bound and optimized reusable plan. +#[derive(Clone)] +pub struct PreparedPlan<'db> { + pub(crate) plan: LogicalPlan, + pub(crate) arena: PlanArena<'db>, + parameter_expressions: Vec, +} + +impl<'db> PreparedPlan<'db> { + pub(crate) fn bind_parameters( + &self, + params: &[(usize, DataValue)], + ) -> Result<(LogicalPlan, ParamArena<'_>), DatabaseError> { + let mut arena = ParamArena::new(&self.arena, &self.parameter_expressions, params)?; + let mut plan = self.plan.clone(); + ParameterBinder { params }.visit_plan(&mut plan)?; + SpecializeIndexRange { arena: &mut arena }.visit_plan(&mut plan)?; + Ok((plan, arena)) + } +} + +impl Database { + /// Prepare a plan with explicit positional parameter types. + /// Parameter values are supplied separately for each execution. + pub fn prepare( + &self, + sql: &str, + params: &[(usize, LogicalType)], + ) -> Result, DatabaseError> { + let statement = crate::binder::parse_statement(sql)?; + let transaction = self + .storage + .transaction_with_isolation(self.transaction_isolation)?; + self.state.prepare_plan(&statement, params, &transaction) + } +} + +impl<'db, S: Storage> DBTransaction<'db, S> { + pub fn prepare( + &self, + sql: &str, + params: &[(usize, LogicalType)], + ) -> Result, DatabaseError> { + let statement = crate::binder::parse_statement(sql)?; + self.state.prepare_plan(&statement, params, &self.inner) + } +} + +impl State { + pub(crate) fn prepare_plan<'a, 'txn>( + &'a self, + statement: &Statement, + params: &[(usize, LogicalType)], + transaction: &S::TransactionType<'txn>, + ) -> Result, DatabaseError> + where + S: 'txn, + { + if matches!( + crate::binder::command_type(statement)?, + crate::binder::CommandType::DDL | crate::binder::CommandType::Analyze + ) { + return Err(DatabaseError::UnsupportedStmt( + "DDL and ANALYZE require ddl/analyze".into(), + )); + } + let (mut plan, mut arena) = self.build_plan(params, transaction, |binder, arena| { + binder.bind(statement, arena) + })?; + plan.output_schema(&mut arena); + let parameter_expressions = arena.parameter_expressions(); + Ok(PreparedPlan { + plan, + arena, + parameter_expressions, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::catalog::{ColumnCatalog, ColumnDesc}; + use crate::expression::range_detacher::Range; + use crate::expression::{BinaryOperator, ScalarExpression}; + use crate::planner::operator::filter::FilterOperator; + use crate::planner::operator::sort::SortField; + use crate::planner::operator::table_scan::TableScanOperator; + use crate::planner::{Childrens, LogicalPlan, TableArenaCell}; + use crate::types::index::{IndexInfo, IndexMeta, IndexType}; + use crate::types::tuple::TupleLike; + use crate::types::value::DataValue; + use crate::types::LogicalType; + use std::ops::Bound; + + #[test] + fn prepare_caches_scalar_output_schema() -> Result<(), DatabaseError> { + let db = DataBaseBuilder::path(".").build_in_memory()?; + let plan = db.prepare( + "select (($1 * 3 + 7) % 97) + ($1 / 2)", + &[(1, LogicalType::Bigint)], + )?; + for (value, expected) in [(7, 31.5), (23, 87.5)] { + let mut iter = db.execute(&plan, [(1, DataValue::Int64(value))])?; + iter.schema(|schema| { + assert_eq!(schema.len(), 1); + assert_eq!( + schema.iter().next().unwrap().datatype(), + &LogicalType::Double + ); + }); + assert_eq!( + iter.next_tuple(|_, row| row.values.clone())?, + Some(vec![DataValue::Float64(expected.into())]) + ); + assert!(iter.next_tuple(|_, _| ())?.is_none()); + iter.done()?; + } + let mut tx = db.new_transaction()?; + let mut iter = tx.execute(&plan, [(1, DataValue::Int64(7))])?; + assert_eq!( + iter.next_tuple(|_, row| row.values.clone())?, + Some(vec![DataValue::Float64(31.5.into())]) + ); + iter.done()?; + tx.commit()?; + Ok(()) + } + + #[test] + fn specialize_selected_index_and_preserve_plan_metadata() -> Result<(), DatabaseError> { + let table_arena = TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let mut column = ColumnCatalog::new( + "id".into(), + false, + ColumnDesc::new(LogicalType::Integer, None, false, None)?, + ); + column.set_ref_table("t".into(), 1, false); + let column = arena.alloc_column(column); + let col_expr = arena.alloc_expression(ScalarExpression::column_expr(column, 0)); + let meta = arena.alloc_index(IndexMeta { + id: 1, + column_ids: vec![1], + table_name: "t".into(), + pk_ty: LogicalType::Integer, + value_ty: LogicalType::Integer, + name: "pk".into(), + ty: IndexType::PrimaryKey { is_multiple: false }, + }); + let original = Range::Scope { + min: Bound::Unbounded, + max: Bound::Excluded(DataValue::Int32(5)), + }; + let expected = Range::Scope { + min: Bound::Included(DataValue::Int32(3)), + max: Bound::Excluded(DataValue::Int32(5)), + }; + let sort = SortOption::OrderBy { + fields: vec![SortField::new(col_expr, true, false)], + ignore_prefix_len: 0, + }; + for (lookup, has_residual, parameter, should_change) in [ + ( + IndexLookup::Static(original.clone()), + true, + DataValue::Int32(5), + true, + ), + ( + IndexLookup::Static(original.clone()), + false, + DataValue::Int32(5), + false, + ), + (IndexLookup::Probe, true, DataValue::Int32(5), false), + ( + IndexLookup::Static(original.clone()), + true, + DataValue::Int32(i32::MIN), + false, + ), + ( + IndexLookup::Static(original.clone()), + true, + DataValue::Null, + false, + ), + ] { + let left = arena.alloc_expression(ScalarExpression::Constant(parameter)); + let right = arena.alloc_expression(ScalarExpression::Constant(DataValue::Int32(2))); + let boundary = arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::Minus, + left_expr: left, + right_expr: right, + evaluator: None, + ty: LogicalType::Integer, + }); + let predicate = arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::GtEq, + left_expr: col_expr, + right_expr: boundary, + evaluator: None, + ty: LogicalType::Boolean, + }); + let index = IndexInfo { + meta, + lookup: Some(lookup), + residual_predicate: has_residual.then_some(predicate), + sort_option: sort.clone(), + covered_deserializers: None, + cover_mapping: None, + sort_elimination_hint: None, + stream_aggregate_hint: None, + }; + let mut other = index.clone(); + other.meta = arena.alloc_index(IndexMeta { + id: 2, + column_ids: vec![1], + table_name: "t".into(), + pk_ty: LogicalType::Integer, + value_ty: LogicalType::Integer, + name: "other".into(), + ty: IndexType::Normal, + }); + let mut scan = LogicalPlan::new( + Operator::TableScan(TableScanOperator { + table_name: "t".into(), + columns: vec![column], + limit: (None, None), + index_infos: vec![index.clone(), other.clone()], + with_pk: false, + }), + Childrens::None, + ); + scan.physical_option = Some(PhysicalOption::new( + PlanImpl::IndexScan(Box::new(index.clone())), + sort.clone(), + )); + let filter = FilterOperator { + predicate, + having: false, + is_optimized: true, + }; + let mut plan = LogicalPlan::new( + Operator::Filter(filter.clone()), + Childrens::Only(Box::new(scan)), + ); + let mut expected_index = index.clone(); + if should_change { + expected_index.lookup = Some(IndexLookup::Static(expected.clone())); + } + SpecializeIndexRange { arena: &mut arena }.visit_plan(&mut plan)?; + assert_eq!(plan.operator, Operator::Filter(filter)); + let child = plan.childrens.pop_only(); + let option = child.physical_option.unwrap(); + assert_eq!( + option.plan, + PlanImpl::IndexScan(Box::new(expected_index.clone())) + ); + assert_eq!(option.sort_option(), &sort); + let Operator::TableScan(scan) = child.operator else { + panic!("expected scan") + }; + assert_eq!(scan.index_infos[0], expected_index); + assert_eq!(scan.index_infos[1], other); + assert_eq!( + arena.expression(boundary), + &ScalarExpression::Binary { + op: BinaryOperator::Minus, + left_expr: left, + right_expr: right, + evaluator: None, + ty: LogicalType::Integer, + } + ); + } + Ok(()) + } + + #[test] + fn parameter_order_dependent_predicates_remain_correct() -> Result<(), DatabaseError> { + fn has_filter(plan: &LogicalPlan) -> bool { + matches!(plan.operator, Operator::Filter(_)) || plan.childrens.iter().any(has_filter) + } + + let mut db = DataBaseBuilder::path(".") + .histogram_buckets(2) + .build_in_memory()?; + db.ddl("create table parameter_order(id int primary key)")?; + db.run("insert into parameter_order values(1),(2),(3),(4)")? + .done()?; + db.analyze("parameter_order")?; + + let lower_bounds = db.prepare( + "select id from parameter_order where id >= $1 and id >= $2 order by id", + &[(1, LogicalType::Integer), (2, LogicalType::Integer)], + )?; + assert!(has_filter(&lower_bounds.plan)); + let mut iter = db.execute( + &lower_bounds, + [(1, DataValue::Int32(3)), (2, DataValue::Int32(1))], + )?; + let mut rows = Vec::new(); + while iter + .next_tuple(|_, row| rows.push(row.value_at(0).clone()))? + .is_some() + {} + iter.done()?; + assert_eq!(rows, vec![DataValue::Int32(3), DataValue::Int32(4)]); + + let equalities = db.prepare( + "select id from parameter_order where id = $1 and id = $2", + &[(1, LogicalType::Integer), (2, LogicalType::Integer)], + )?; + let mut iter = db.execute( + &equalities, + [(1, DataValue::Int32(2)), (2, DataValue::Int32(2))], + )?; + assert_eq!( + iter.next_tuple(|_, row| row.value_at(0).clone())?, + Some(DataValue::Int32(2)) + ); + assert!(iter.next_tuple(|_, _| ())?.is_none()); + iter.done()?; + + let mut iter = db.execute( + &equalities, + [(1, DataValue::Int32(2)), (2, DataValue::Int32(3))], + )?; + assert!(iter.next_tuple(|_, _| ())?.is_none()); + iter.done()?; + Ok(()) + } + + #[test] + fn repeated_execution_and_null_do_not_stale_parameters() -> Result<(), DatabaseError> { + let db = DataBaseBuilder::path(".").build_in_memory()?; + let plan = db.prepare( + "values ($1 + $2), ($2)", + &[(1, LogicalType::Integer), (2, LogicalType::Integer)], + )?; + for (value, expected) in [ + (DataValue::Null, DataValue::Null), + (DataValue::Int32(3), DataValue::Int32(13)), + ] { + let mut iter = db.execute(&plan, [(1, DataValue::Int32(10)), (2, value.clone())])?; + assert_eq!( + iter.next_tuple(|_, row| row.value_at(0).clone())?, + Some(expected) + ); + assert_eq!( + iter.next_tuple(|_, row| row.value_at(0).clone())?, + Some(value) + ); + assert!(iter.next_tuple(|_, _| ())?.is_none()); + iter.done()?; + } + + assert!(matches!( + db.execute(&plan, [(1, DataValue::Int32(10))]), + Err(DatabaseError::ParametersNotFound { name, .. }) if name == "$2" + )); + let mut iter = db.execute(&plan, [(1, DataValue::Int32(10)), (2, DataValue::Int32(4))])?; + assert_eq!( + iter.next_tuple(|_, row| row.value_at(0).clone())?, + Some(DataValue::Int32(14)) + ); + iter.done()?; + + Ok(()) + } + + #[test] + fn composite_index_ranges_bind_for_each_execution() -> Result<(), DatabaseError> { + fn index_range(plan: &LogicalPlan) -> Option<&Range> { + plan.physical_option + .as_ref() + .and_then(|option| match &option.plan { + PlanImpl::IndexScan(info) => match &info.lookup { + Some(IndexLookup::Static(range)) => Some(range), + _ => None, + }, + _ => None, + }) + .or_else(|| plan.childrens.iter().find_map(index_range)) + } + + let mut db = DataBaseBuilder::path(".") + .histogram_buckets(2) + .build_in_memory()?; + db.ddl("create table t(w int, k int, primary key(w,k))")?; + db.run("insert into t values(1,1),(1,2),(1,3),(2,1),(2,2),(2,4)")? + .done()?; + db.analyze("t")?; + let plan = db.prepare( + "select k from t where w=$1 and k >= $2 and k < $3 order by k", + &[ + (1, LogicalType::Integer), + (2, LogicalType::Integer), + (3, LogicalType::Integer), + ], + )?; + assert!(index_range(&plan.plan).is_some()); + let mut iter = db.execute( + &plan, + [ + (1, DataValue::Int32(2)), + (2, DataValue::Int32(2)), + (3, DataValue::Int32(5)), + ], + )?; + let mut rows = Vec::new(); + while iter + .next_tuple(|_, row| rows.push(row.value_at(0).clone()))? + .is_some() + {} + iter.done()?; + assert_eq!(rows, vec![DataValue::Int32(2), DataValue::Int32(4)]); + let mut iter = db.execute( + &plan, + [ + (1, DataValue::Int32(1)), + (2, DataValue::Int32(1)), + (3, DataValue::Int32(2)), + ], + )?; + assert_eq!( + iter.next_tuple(|_, row| row.value_at(0).clone())?, + Some(DataValue::Int32(1)) + ); + assert!(iter.next_tuple(|_, _| ())?.is_none()); + iter.done()?; + + // The lower bound is computed only after binding; it must still be + // intersected with the prepared upper bound instead of scanning the + // whole equality prefix and filtering rows afterwards. + let plan = db.prepare( + "select k from t where w=$1 and k<$2 and k>=($3-20)", + &[ + (1, LogicalType::Integer), + (2, LogicalType::Integer), + (3, LogicalType::Integer), + ], + )?; + let (bound, _) = plan.bind_parameters(&[ + (1, DataValue::Int32(2)), + (2, DataValue::Int32(5)), + (3, DataValue::Int32(22)), + ])?; + assert_eq!( + index_range(&bound), + Some(&Range::Scope { + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(2), + DataValue::Int32(2), + ])), + max: Bound::Excluded(DataValue::Tuple(vec![ + DataValue::Int32(2), + DataValue::Int32(5), + ])), + }) + ); + Ok(()) + } +} diff --git a/src/execution/ddl/add_column.rs b/src/execution/ddl/add_column.rs index fc77db96..758025ee 100644 --- a/src/execution/ddl/add_column.rs +++ b/src/execution/ddl/add_column.rs @@ -19,6 +19,7 @@ use crate::execution::{ }; use crate::iter_ext::Itertools; use crate::planner::operator::alter_table::add_column::AddColumnOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::index::{Index, IndexType}; use crate::types::tuple_builder::TupleBuilder; @@ -40,7 +41,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for AddColumn { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -53,7 +54,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for AddColumn { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let table_cache = arena.table_cache(); let Some(AddColumnOperator { @@ -79,7 +80,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for AddColumn { }; if column_exists { if if_not_exists { - TupleBuilder::build_result_into(arena.result_tuple_mut(), "1".to_string()); + arena.produce_tuple(TupleBuilder::build_result("1".to_string())); arena.resume(); return Ok(()); } @@ -143,14 +144,18 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for AddColumn { default_for_index.as_ref(), tuple.pk.as_ref(), ) { - let index = Index::new(*unique_index_id, value, IndexType::Unique); + let index = Index::new( + *unique_index_id, + std::slice::from_ref(value), + IndexType::Unique, + ); transaction.add_index(table_codec, &table_name, index, tuple_id)?; } Ok(()) }, )?; - TupleBuilder::build_result_into(arena.result_tuple_mut(), "1".to_string()); + arena.produce_tuple(TupleBuilder::build_result("1".to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/change_column.rs b/src/execution/ddl/change_column.rs index 40265bc8..350024f7 100644 --- a/src/execution/ddl/change_column.rs +++ b/src/execution/ddl/change_column.rs @@ -19,6 +19,7 @@ use crate::execution::{ }; use crate::iter_ext::Itertools; use crate::planner::operator::alter_table::change_column::{ChangeColumnOperator, NotNullChange}; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -38,7 +39,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for ChangeColumn { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -51,7 +52,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ChangeColumn { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let table_cache = arena.table_cache(); let Some(ChangeColumnOperator { @@ -188,7 +189,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ChangeColumn { }; arena.push_ddl_apply(apply); - TupleBuilder::build_result_into(arena.result_tuple_mut(), format!("{table_name}")); + arena.produce_tuple(TupleBuilder::build_result(format!("{table_name}"))); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/create_index.rs b/src/execution/ddl/create_index.rs index 3440581f..9ad43081 100644 --- a/src/execution/ddl/create_index.rs +++ b/src/execution/ddl/create_index.rs @@ -14,12 +14,13 @@ use crate::errors::DatabaseError; use crate::execution::{ - build_read, with_projection_tmp_value, DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, - ExecutorNode, WriteExecutor, + build_read, DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, + WriteExecutor, }; use crate::expression::ScalarExpression; use crate::planner::operator::create_index::CreateIndexOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::index::Index; use crate::types::tuple::Schema; @@ -50,7 +51,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for CreateIndex { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -71,7 +72,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateIndex { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(CreateIndexOperator { table_name, @@ -137,15 +138,16 @@ 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, plan_arena, None, &column_exprs, |arena, value| { + arena.rewrite(&column_exprs, plan_arena, None)?; + { let mut state = arena.local_state(plan_arena); - let (transaction, table_codec) = state.transaction_codec_mut(); - let index = Index::new(index_id, &value, ty); - transaction.add_index(table_codec, table_name.as_ref(), index, &tuple_pk) - })?; + let (values, transaction, table_codec) = state.index_values_transaction_codec_mut(); + let index = Index::new(index_id, values, ty); + transaction.add_index(table_codec, table_name.as_ref(), index, &tuple_pk)?; + } } - TupleBuilder::build_result_into(arena.result_tuple_mut(), "1".to_string()); + arena.produce_tuple(TupleBuilder::build_result("1".to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/create_table.rs b/src/execution/ddl/create_table.rs index fc275e7f..b9cab332 100644 --- a/src/execution/ddl/create_table.rs +++ b/src/execution/ddl/create_table.rs @@ -17,6 +17,7 @@ use crate::execution::{ DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::create_table::CreateTableOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for CreateTable { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateTable { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(CreateTableOperator { table_name, @@ -73,7 +74,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateTable { arena.push_ddl_apply(DDLApply::upsert_table(table, false)); } - TupleBuilder::build_result_into(arena.result_tuple_mut(), format!("{table_name}")); + arena.produce_tuple(TupleBuilder::build_result(format!("{table_name}"))); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/create_view.rs b/src/execution/ddl/create_view.rs index c676c83c..22c55f1f 100644 --- a/src/execution/ddl/create_view.rs +++ b/src/execution/ddl/create_view.rs @@ -17,6 +17,7 @@ use crate::execution::{ DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::create_view::CreateViewOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for CreateView { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateView { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(CreateViewOperator { view, or_replace }) = self.op.take() else { arena.finish(); @@ -60,7 +61,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CreateView { let view = transaction.create_view(table_codec, plan_arena, view, or_replace)?; arena.push_ddl_apply(DDLApply::upsert_view(view)); - TupleBuilder::build_result_into(arena.result_tuple_mut(), view_name); + arena.produce_tuple(TupleBuilder::build_result(view_name)); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/drop_column.rs b/src/execution/ddl/drop_column.rs index d3b35473..e2b47cb7 100644 --- a/src/execution/ddl/drop_column.rs +++ b/src/execution/ddl/drop_column.rs @@ -19,6 +19,7 @@ use crate::execution::{ }; use crate::iter_ext::Itertools; use crate::planner::operator::alter_table::drop_column::DropColumnOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -38,7 +39,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for DropColumn { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -51,7 +52,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropColumn { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let table_cache = arena.table_cache(); let Some(DropColumnOperator { @@ -124,7 +125,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropColumn { arena.push_ddl_apply(DDLApply::upsert_table(table, true)); } - TupleBuilder::build_result_into(arena.result_tuple_mut(), "1".to_string()); + arena.produce_tuple(TupleBuilder::build_result("1".to_string())); arena.resume(); Ok(()) } else if !if_exists { diff --git a/src/execution/ddl/drop_index.rs b/src/execution/ddl/drop_index.rs index e84062bd..8f34aab5 100644 --- a/src/execution/ddl/drop_index.rs +++ b/src/execution/ddl/drop_index.rs @@ -17,6 +17,7 @@ use crate::execution::{ DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::drop_index::DropIndexOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for DropIndex { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropIndex { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(DropIndexOperator { table_name, @@ -79,7 +80,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropIndex { }); } - TupleBuilder::build_result_into(arena.result_tuple_mut(), index_name.to_string()); + arena.produce_tuple(TupleBuilder::build_result(index_name.to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/drop_table.rs b/src/execution/ddl/drop_table.rs index 4e975249..b3470f54 100644 --- a/src/execution/ddl/drop_table.rs +++ b/src/execution/ddl/drop_table.rs @@ -17,6 +17,7 @@ use crate::execution::{ DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::drop_table::DropTableOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for DropTable { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropTable { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(DropTableOperator { table_name, @@ -67,7 +68,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropTable { }); } - TupleBuilder::build_result_into(arena.result_tuple_mut(), format!("{table_name}")); + arena.produce_tuple(TupleBuilder::build_result(format!("{table_name}"))); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/drop_view.rs b/src/execution/ddl/drop_view.rs index 5f34f2ce..891d31d8 100644 --- a/src/execution/ddl/drop_view.rs +++ b/src/execution/ddl/drop_view.rs @@ -17,6 +17,7 @@ use crate::execution::{ DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::drop_view::DropViewOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for DropView { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropView { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - _: &mut crate::planner::PlanArena<'a>, + _: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(DropViewOperator { view_name, @@ -67,7 +68,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for DropView { }); } - TupleBuilder::build_result_into(arena.result_tuple_mut(), format!("{view_name}")); + arena.produce_tuple(TupleBuilder::build_result(format!("{view_name}"))); arena.resume(); Ok(()) } diff --git a/src/execution/ddl/truncate.rs b/src/execution/ddl/truncate.rs index 49277bba..072b0b2e 100644 --- a/src/execution/ddl/truncate.rs +++ b/src/execution/ddl/truncate.rs @@ -17,6 +17,7 @@ use crate::execution::{ ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::truncate::TruncateOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -36,7 +37,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for Truncate { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -49,7 +50,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Truncate { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(TruncateOperator { table_name }) = self.op.take() else { arena.finish(); @@ -59,7 +60,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Truncate { let (transaction, table_codec) = state.transaction_codec_mut(); transaction.drop_data(table_codec, &table_name)?; - TupleBuilder::build_result_into(arena.result_tuple_mut(), format!("{table_name}")); + arena.produce_tuple(TupleBuilder::build_result(format!("{table_name}"))); arena.resume(); Ok(()) } diff --git a/src/execution/dml/analyze.rs b/src/execution/dml/analyze.rs index 4d0b3f48..116f355f 100644 --- a/src/execution/dml/analyze.rs +++ b/src/execution/dml/analyze.rs @@ -15,8 +15,8 @@ use crate::catalog::TableName; use crate::errors::DatabaseError; use crate::execution::{ - build_read, with_projection_tmp_value, DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, - ExecutorNode, WriteExecutor, + build_read, DDLApply, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, + WriteExecutor, }; use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; @@ -24,6 +24,7 @@ use crate::optimizer::core::histogram::{HistogramBuilder, ANALYZE_STATISTICS_REL use crate::optimizer::core::statistics_meta::StatisticsMeta; use crate::planner::operator::analyze::AnalyzeOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::{table_codec::TableCodec, Transaction}; use crate::types::index::IndexId; use crate::types::value::{DataValue, Utf8Type}; @@ -66,7 +67,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for Analyze { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -86,7 +87,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Analyze { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(input) = self.input.take() else { arena.finish(); @@ -102,10 +103,13 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Analyze { .indexes() .map(|index| { let index = plan_arena.index(*index); + let index_id = index.id; + let builder = HistogramBuilder::new(index, ANALYZE_STATISTICS_RELATIVE_ERROR)?; + let exprs = index.column_exprs(table).collect::, _>>()?; Ok(State { - index_id: index.id, - exprs: index.column_exprs(table, plan_arena)?, - builder: HistogramBuilder::new(index, ANALYZE_STATISTICS_RELATIVE_ERROR)?, + index_id, + exprs, + builder, histogram_buckets: self.histogram_buckets, }) }) @@ -113,10 +117,15 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Analyze { }; while arena.next_tuple(input, plan_arena)? { + let tuple = arena.materialize_tuple(); for State { exprs, builder, .. } in builders.iter_mut() { - with_projection_tmp_value(arena, plan_arena, None, exprs, |_, value| { - builder.append(value) - })?; + arena.rewrite(exprs, plan_arena, Some(&tuple))?; + let key = arena.materialize_tuple().values; + let value = match <[DataValue; 1]>::try_from(key) { + Ok([value]) => value, + Err(key) => DataValue::Tuple(key), + }; + builder.append(value)?; } } let mut state = arena.local_state(plan_arena); @@ -130,10 +139,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Analyze { plan_arena, )?; - let output = arena.result_tuple_mut(); - output.pk = None; - output.values = values; - arena.resume(); + arena.produce_tuple(crate::types::tuple::Tuple::new(None, values)); Ok(()) } } @@ -152,7 +158,7 @@ impl Analyze { applies: &mut Vec, transaction: &mut U, table_codec: &mut TableCodec, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { let mut values = Vec::with_capacity(builders.len()); diff --git a/src/execution/dml/copy_from_file.rs b/src/execution/dml/copy_from_file.rs index 6e1ca075..72684157 100644 --- a/src/execution/dml/copy_from_file.rs +++ b/src/execution/dml/copy_from_file.rs @@ -19,6 +19,7 @@ use crate::execution::{ }; use crate::iter_ext::Itertools; use crate::planner::operator::copy_from_file::CopyFromFileOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; use std::fs::File; @@ -40,7 +41,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for CopyFromFile { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -53,7 +54,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CopyFromFile { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(op) = self.op.take() else { arena.finish(); @@ -111,7 +112,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CopyFromFile { size += 1; } - TupleBuilder::build_result_into(arena.result_tuple_mut(), size.to_string()); + arena.produce_tuple(TupleBuilder::build_result(size.to_string())); arena.resume(); Ok(()) } @@ -125,6 +126,7 @@ mod tests { use crate::db::{CatalogKind, DataBaseBuilder}; use crate::errors::DatabaseError; use crate::storage::Storage; + use crate::types::tuple::TupleLike; use crate::types::CharLengthUnits; use crate::types::LogicalType; use std::io::Write; @@ -196,7 +198,7 @@ mod tests { let result = executor .next_tuple()? .expect("copy from file should yield once"); - assert_eq!(result.values[0].to_string(), "2"); + assert_eq!(result.value_at(0).to_string(), "2"); Ok(()) } diff --git a/src/execution/dml/copy_to_file.rs b/src/execution/dml/copy_to_file.rs index a79b40b2..d9402c7b 100644 --- a/src/execution/dml/copy_to_file.rs +++ b/src/execution/dml/copy_to_file.rs @@ -20,6 +20,7 @@ use crate::execution::{ use crate::iter_ext::Itertools; use crate::planner::operator::copy_to_file::CopyToFileOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple_builder::TupleBuilder; @@ -47,7 +48,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for CopyToFile { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -96,7 +97,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CopyToFile { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(input) = self.input.take() else { arena.finish(); @@ -105,7 +106,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CopyToFile { let mut writer = self.create_writer()?; while arena.next_tuple(input, plan_arena)? { - let tuple = arena.result_tuple(); + let tuple = arena.materialize_tuple(); writer.write_record( tuple .values @@ -121,7 +122,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for CopyToFile { } else { format!("{} [{}]", self.op, self.column_names.iter().join(", ")) }; - TupleBuilder::build_result_into(arena.result_tuple_mut(), message); + arena.produce_tuple(TupleBuilder::build_result(message)); arena.resume(); Ok(()) } @@ -135,6 +136,7 @@ mod tests { use crate::errors::DatabaseError; use crate::planner::operator::table_scan::TableScanOperator; use crate::storage::Storage; + use crate::types::tuple::TupleLike; use tempfile::TempDir; #[test] @@ -206,7 +208,7 @@ mod tests { let record3 = records.next().unwrap()?; assert_eq!(record3, vec!["3", "2.1", "Kite"]); - assert_eq!(tuple.values[0].to_string(), format!("{op} [a, b, c]")); + assert_eq!(tuple.value_at(0).to_string(), format!("{op} [a, b, c]")); Ok(()) } } diff --git a/src/execution/dml/delete.rs b/src/execution/dml/delete.rs index 158002c0..e16680a6 100644 --- a/src/execution/dml/delete.rs +++ b/src/execution/dml/delete.rs @@ -15,11 +15,11 @@ use crate::catalog::TableName; use crate::errors::DatabaseError; use crate::execution::{ - build_read, with_projection_tmp_value, ExecArena, ExecId, ExecNode, ExecutionContext, - ExecutorNode, WriteExecutor, + build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::planner::operator::delete::DeleteOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::index::Index; use crate::types::tuple_builder::TupleBuilder; @@ -46,7 +46,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for Delete { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -66,7 +66,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Delete { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(input) = self.input.take() else { arena.finish(); @@ -85,7 +85,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Delete { Ok(( index_meta.id, index_meta.ty, - index_meta.column_exprs(table, plan_arena)?, + index_meta + .column_exprs(table) + .collect::, _>>()?, )) }) .collect::, DatabaseError>>()? @@ -97,17 +99,13 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Delete { continue; }; + let tuple = arena.materialize_tuple(); for (index_id, index_ty, exprs) in index_templates.iter() { - 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( - table_codec, - &self.table_name, - &Index::new(*index_id, &value, *index_ty), - &tuple_id, - ) - })?; + arena.rewrite(exprs, plan_arena, Some(&tuple))?; + let mut state = arena.local_state(plan_arena); + let (values, transaction, table_codec) = state.index_values_transaction_codec_mut(); + let index = Index::new(*index_id, values, *index_ty); + transaction.del_index(table_codec, &self.table_name, &index, &tuple_id)?; } let mut state = arena.local_state(plan_arena); @@ -116,7 +114,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Delete { deleted_count += 1; } - TupleBuilder::build_result_into(arena.result_tuple_mut(), deleted_count.to_string()); + arena.produce_tuple(TupleBuilder::build_result(deleted_count.to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/dml/insert.rs b/src/execution/dml/insert.rs index 4e91952a..6b781d3a 100644 --- a/src/execution/dml/insert.rs +++ b/src/execution/dml/insert.rs @@ -15,12 +15,12 @@ use crate::catalog::TableName; use crate::errors::DatabaseError; use crate::execution::{ - build_read, with_projection_tmp_value, ExecArena, ExecId, ExecNode, ExecutionContext, - ExecutorNode, WriteExecutor, + build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::iter_ext::Itertools; use crate::planner::operator::insert::InsertOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::index::Index; use crate::types::tuple::{Schema, Tuple}; @@ -66,7 +66,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for Insert { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -106,7 +106,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Insert { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(input) = self.input.take() else { arena.finish(); @@ -136,7 +136,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Insert { while arena.next_tuple(input, plan_arena)? { let mut tuple_map = HashMap::with_capacity(self.input_schema.len()); - for (i, value) in arena.result_tuple_mut().values.drain(..).enumerate() { + for (i, value) in arena.materialize_tuple().values.into_iter().enumerate() { let column = plan_arena.column(self.input_schema[i]); tuple_map.insert(Self::column_key(column, self.is_mapping_by_name), value); } @@ -168,18 +168,12 @@ 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, - 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) - }, - )?; + arena.rewrite(exprs, plan_arena, Some(&tuple))?; + let mut state = arena.local_state(plan_arena); + let (values, transaction, table_codec) = + state.index_values_transaction_codec_mut(); + let index = Index::new(index_meta.id, values, 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(); @@ -193,11 +187,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Insert { inserted_count += 1; } - TupleBuilder::build_result_into(arena.result_tuple_mut(), inserted_count.to_string()); + arena.produce_tuple(TupleBuilder::build_result(inserted_count.to_string())); arena.resume(); Ok(()) } else { - TupleBuilder::build_result_into(arena.result_tuple_mut(), "0".to_string()); + arena.produce_tuple(TupleBuilder::build_result("0".to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/dml/update.rs b/src/execution/dml/update.rs index 963b84fc..dfa969cd 100644 --- a/src/execution/dml/update.rs +++ b/src/execution/dml/update.rs @@ -15,21 +15,18 @@ use crate::catalog::{ColumnRef, TableName}; use crate::errors::DatabaseError; use crate::execution::{ - build_read, with_projection_tmp_value, ExecArena, ExecId, ExecNode, ExecutionContext, - ExecutorNode, WriteExecutor, + build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, WriteExecutor, }; use crate::iter_ext::Itertools; use crate::planner::operator::update::UpdateOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::index::{Index, IndexMeta, IndexType}; use crate::types::tuple::{Schema, Tuple}; use crate::types::tuple_builder::TupleBuilder; use crate::types::ColumnId; -use std::{ - collections::{HashMap, HashSet}, - mem, -}; +use std::collections::{HashMap, HashSet}; pub struct Update { table_name: TableName, @@ -65,7 +62,7 @@ impl<'a, T: Transaction + 'a> WriteExecutor<'a, T> for Update { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -104,7 +101,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(input) = self.input.take() else { arena.finish(); @@ -149,7 +146,8 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { while arena.next_tuple(input, plan_arena)? { let mut is_overwrite = true; - let Some(old_pk) = arena.result_tuple().pk.clone() else { + let mut tuple = arena.materialize_tuple(); + let Some(old_pk) = tuple.pk.clone() else { continue; }; @@ -166,10 +164,9 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { continue; } - with_projection_tmp_value(arena, plan_arena, None, exprs, |_, value| { - old_index_values.push((index_offset, value)); - Ok(()) - })?; + arena.rewrite(exprs, plan_arena, Some(&tuple))?; + let values = arena.materialize_tuple().values; + old_index_values.push((index_offset, values)); } for (i, column) in self.input_schema.iter().enumerate() { let Some(column_id) = plan_arena.column(*column).id() else { @@ -178,16 +175,14 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { if let Some(expr) = exprs_map.get(&column_id) { let value = plan_arena .expression(*expr) - .eval(plan_arena, Some(arena.result_tuple()))?; - arena.result_tuple_mut().values[i] = value; + .eval(plan_arena, Some(&tuple))?; + tuple.values[i] = value.into_owned(); } } - let new_pk = Tuple::primary_projection( - table_snapshot.primary_key_indices, - &arena.result_tuple().values, - ); - arena.result_tuple_mut().pk = Some(new_pk.clone()); + let new_pk = + Tuple::primary_projection(table_snapshot.primary_key_indices, &tuple.values); + tuple.pk = Some(new_pk.clone()); let primary_key_changed = new_pk != old_pk; if primary_key_changed { @@ -202,26 +197,20 @@ 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, plan_arena, None, exprs, |arena, value| { - if !primary_key_changed && old_value == value { - return Ok(()); - } - - let mut state = arena.local_state(plan_arena); - let (transaction, table_codec) = state.transaction_codec_mut(); - let old_index = Index::new(index_id, &old_value, index_ty); - transaction.del_index( - table_codec, - &self.table_name, - &old_index, - &old_pk, - )?; - let new_index = Index::new(index_id, &value, index_ty); - transaction.add_index(table_codec, &self.table_name, new_index, &new_pk) - })?; + arena.rewrite(exprs, plan_arena, Some(&tuple))?; + let mut state = arena.local_state(plan_arena); + let (values, transaction, table_codec) = + state.index_values_transaction_codec_mut(); + if !primary_key_changed && old_value == values { + continue; + } + + let old_index = Index::new(index_id, &old_value, index_ty); + let new_index = Index::new(index_id, values, index_ty); + transaction.del_index(table_codec, &self.table_name, &old_index, &old_pk)?; + transaction.add_index(table_codec, &self.table_name, new_index, &new_pk)?; } - let tuple = mem::take(arena.result_tuple_mut()); let mut state = arena.local_state(plan_arena); let (transaction, table_codec) = state.transaction_codec_mut(); transaction.append_tuple( @@ -234,11 +223,11 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Update { updated_count += 1; } - TupleBuilder::build_result_into(arena.result_tuple_mut(), updated_count.to_string()); + arena.produce_tuple(TupleBuilder::build_result(updated_count.to_string())); arena.resume(); Ok(()) } else { - TupleBuilder::build_result_into(arena.result_tuple_mut(), "0".to_string()); + arena.produce_tuple(TupleBuilder::build_result("0".to_string())); arena.resume(); Ok(()) } diff --git a/src/execution/dql/aggregate/hash_agg.rs b/src/execution/dql/aggregate/hash_agg.rs index 54bf56e1..cb66f900 100644 --- a/src/execution/dql/aggregate/hash_agg.rs +++ b/src/execution/dql/aggregate/hash_agg.rs @@ -20,6 +20,7 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::aggregate::AggregateOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -48,7 +49,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for HashAggExecutor { input, ): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -66,7 +67,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashAggExecutor { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.output.is_none() { let mut group_hash_accs: HashMap, Vec>> = @@ -77,7 +78,12 @@ 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(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); + group_keys.push( + plan_arena + .expression(*expr) + .eval(plan_arena, Some(tuple))? + .into_owned(), + ); } if let Some(accs) = group_hash_accs.get_mut(group_keys.as_slice()) { @@ -119,6 +125,7 @@ mod test { use crate::planner::operator::aggregate::AggregateOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::Storage; @@ -146,8 +153,8 @@ mod test { ]; let input = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![ DataValue::Int32(0), DataValue::Int32(2), @@ -168,9 +175,10 @@ mod test { DataValue::Int32(2), DataValue::Int32(3), ], - ], - schema_ref: t1_schema.clone(), - }), + ]), + 4, + t1_schema.clone(), + )), Childrens::None, ); let groupby_expr = diff --git a/src/execution/dql/aggregate/mod.rs b/src/execution/dql/aggregate/mod.rs index 3c27e46d..96ebdf52 100644 --- a/src/execution/dql/aggregate/mod.rs +++ b/src/execution/dql/aggregate/mod.rs @@ -29,8 +29,9 @@ 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::planner::ExprRef; +use crate::planner::MetaArena; +use crate::types::tuple::{Tuple, TupleLike}; use crate::types::value::DataValue; use std::borrow::Cow; @@ -70,7 +71,7 @@ pub(crate) fn create_accumulator( #[inline] pub(crate) fn create_accumulators( exprs: &[ExprRef], - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result>, DatabaseError> { exprs .iter() @@ -89,8 +90,8 @@ pub(crate) fn create_accumulators( pub(crate) fn update_accumulators( accs: &mut [Box], agg_calls: &[ExprRef], - tuple: &Tuple, - arena: &PlanArena<'_>, + tuple: &dyn TupleLike, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for (acc, expr) in accs.iter_mut().zip(agg_calls.iter()) { let ScalarExpression::AggCall { args, .. } = arena.expression(*expr) else { diff --git a/src/execution/dql/aggregate/simple_agg.rs b/src/execution/dql/aggregate/simple_agg.rs index 8b6202aa..e07be3c8 100644 --- a/src/execution/dql/aggregate/simple_agg.rs +++ b/src/execution/dql/aggregate/simple_agg.rs @@ -19,6 +19,7 @@ use crate::execution::{ }; use crate::expression::ScalarExpression; use crate::planner::operator::aggregate::AggregateOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; pub struct SimpleAggExecutor { @@ -33,7 +34,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for SimpleAggExecutor { fn into_executor( (AggregateOperator { agg_calls, .. }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -50,7 +51,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for SimpleAggExecutor { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.returned { arena.finish(); diff --git a/src/execution/dql/aggregate/stream_agg.rs b/src/execution/dql/aggregate/stream_agg.rs index d120d705..2b8ebb7f 100644 --- a/src/execution/dql/aggregate/stream_agg.rs +++ b/src/execution/dql/aggregate/stream_agg.rs @@ -20,6 +20,7 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::aggregate::AggregateOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -47,7 +48,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for StreamAggExecutor { input, ): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -66,7 +67,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for StreamAggExecutor { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { if !arena.next_tuple(self.input, plan_arena)? { @@ -86,7 +87,12 @@ 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(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); + group_keys.push( + plan_arena + .expression(*expr) + .eval(plan_arena, Some(tuple))? + .into_owned(), + ); } match &mut self.group_keys { @@ -123,6 +129,7 @@ mod tests { use crate::planner::operator::aggregate::AggregateOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::memory::MemoryStorage; use crate::storage::Storage; @@ -140,16 +147,17 @@ mod tests { }) .to_vec(); let input = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![1.into(), 10.into()], vec![1.into(), 20.into()], vec![2.into(), 5.into()], vec![2.into(), DataValue::Null], vec![2.into(), 7.into()], - ], - schema_ref: columns.clone(), - }), + ]), + 5, + columns.clone(), + )), Childrens::None, ); let group = plan_arena.alloc_expression(ScalarExpression::column_expr(columns[0], 0)); @@ -205,10 +213,7 @@ mod tests { ColumnDesc::new(LogicalType::Integer, None, false, None)?, )); let input = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: Vec::new(), - schema_ref: vec![column], - }), + Operator::Values(ValuesOperator::new(Vec::new(), 0, vec![column])), Childrens::None, ); let operator = AggregateOperator { diff --git a/src/execution/dql/aggregate/stream_distinct.rs b/src/execution/dql/aggregate/stream_distinct.rs index 7b1d3148..0152418e 100644 --- a/src/execution/dql/aggregate/stream_distinct.rs +++ b/src/execution/dql/aggregate/stream_distinct.rs @@ -16,18 +16,16 @@ use crate::errors::DatabaseError; use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; -use crate::iter_ext::Itertools; use crate::planner::operator::aggregate::AggregateOperator; +use crate::planner::MetaArena; 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, input: ExecId, - last_keys: Option>, - scratch: Tuple, + last_keys: Option, } impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for StreamDistinctExecutor { @@ -36,7 +34,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for StreamDistinctExecutor { fn into_executor( (op, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -45,7 +43,6 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for StreamDistinctExecutor { groupby_exprs: op.groupby_exprs, input, last_keys: None, - scratch: Tuple::default(), })) } } @@ -54,29 +51,29 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for StreamDistinctExecutor { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { if !arena.next_tuple(self.input, plan_arena)? { - arena.finish(); + if let Some(last_keys) = self.last_keys.take() { + drop(std::mem::replace(arena.result_tuple_mut(), last_keys)); + arena.resume(); + } else { + arena.finish(); + } return Ok(()); } - std::mem::swap(&mut self.scratch, arena.result_tuple_mut()); - let tuple = &self.scratch; - let group_keys = self - .groupby_exprs - .iter() - .map(|expr| plan_arena.expression(*expr).eval(plan_arena, Some(tuple))) - .try_collect()?; - - if self.last_keys.as_ref() != Some(&group_keys) { - self.last_keys = Some(group_keys.clone()); - let output = arena.result_tuple_mut(); - output.pk.clone_from(&tuple.pk); - output.values = group_keys; + arena.rewrite(&self.groupby_exprs, plan_arena, None)?; + + if let Some(last_keys) = &mut self.last_keys { + if last_keys.values == arena.result_tuple().values { + continue; + } + std::mem::swap(last_keys, arena.result_tuple_mut()); arena.resume(); return Ok(()); } + self.last_keys = Some(arena.materialize_tuple()); } } } @@ -95,6 +92,7 @@ mod tests { use crate::planner::operator::aggregate::AggregateOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::{StatisticsMetaCache, Storage, TableCache, ViewCache}; @@ -147,16 +145,17 @@ mod tests { vec![plan_arena.alloc_column(ColumnCatalog::new("c1".to_string(), true, desc))]; let input = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![DataValue::Int32(1)], vec![DataValue::Int32(1)], vec![DataValue::Int32(2)], vec![DataValue::Int32(2)], vec![DataValue::Int32(3)], - ], - schema_ref: schema_ref.clone(), - }), + ]), + 5, + schema_ref.clone(), + )), Childrens::None, ); let agg = AggregateOperator { @@ -203,16 +202,17 @@ mod tests { ]; let input = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![DataValue::Int32(1), DataValue::Int32(1)], vec![DataValue::Int32(1), DataValue::Int32(1)], vec![DataValue::Int32(1), DataValue::Int32(2)], vec![DataValue::Int32(2), DataValue::Int32(1)], vec![DataValue::Int32(2), DataValue::Int32(1)], - ], - schema_ref: schema_ref.clone(), - }), + ]), + 5, + schema_ref.clone(), + )), Childrens::None, ); let agg = AggregateOperator { diff --git a/src/execution/dql/describe.rs b/src/execution/dql/describe.rs index ba0e227b..4aff2d41 100644 --- a/src/execution/dql/describe.rs +++ b/src/execution/dql/describe.rs @@ -16,6 +16,7 @@ use crate::catalog::{ColumnCatalog, ColumnRef, TableName}; use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::planner::operator::describe::DescribeOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; @@ -61,7 +62,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Describe { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -74,7 +75,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Describe { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.columns.is_none() { let table = arena @@ -109,7 +110,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Describe { } } -fn describe_default(column: &ColumnCatalog, arena: &crate::planner::PlanArena) -> String { +fn describe_default(column: &ColumnCatalog, arena: &(dyn MetaArena + '_)) -> String { column .desc() .default diff --git a/src/execution/dql/dummy.rs b/src/execution/dql/dummy.rs index 3c1d4445..da3e2c73 100644 --- a/src/execution/dql/dummy.rs +++ b/src/execution/dql/dummy.rs @@ -14,6 +14,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; @@ -35,7 +36,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Dummy { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -48,7 +49,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Dummy { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - _: &mut crate::planner::PlanArena<'a>, + _: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let Some(row) = self.row.take() else { arena.finish(); diff --git a/src/execution/dql/explain.rs b/src/execution/dql/explain.rs index feebadfe..a97bb8e5 100644 --- a/src/execution/dql/explain.rs +++ b/src/execution/dql/explain.rs @@ -15,6 +15,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; @@ -39,7 +40,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Explain { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -52,7 +53,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Explain { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.emitted { arena.finish(); diff --git a/src/execution/dql/external_sort.rs b/src/execution/dql/external_sort.rs index 675060eb..2e7fe6df 100644 --- a/src/execution/dql/external_sort.rs +++ b/src/execution/dql/external_sort.rs @@ -20,6 +20,7 @@ use crate::execution::{ }; use crate::planner::operator::sort::{SortField, SortOperator}; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use std::fs::File; use std::io::BufReader; @@ -54,7 +55,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for ExternalSort { fn into_executor( (SortOperator { sort_fields }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -71,7 +72,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ExternalSort { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { if let Some(rows) = &mut self.rows { @@ -92,7 +93,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ExternalSort { let mut rows = SpillVec::new().on_flush(move |rows| sort_segment(sort_fields, rows)); let mut runs = Vec::new(); while arena.next_tuple(self.input, plan_arena)? { - let tuple = mem::take(arena.result_tuple_mut()); + let tuple = arena.materialize_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 cbf80648..d1cd3c15 100644 --- a/src/execution/dql/filter.rs +++ b/src/execution/dql/filter.rs @@ -17,6 +17,7 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::filter::FilterOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; pub struct Filter { @@ -30,7 +31,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Filter { fn into_executor( (FilterOperator { predicate, .. }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -43,7 +44,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Filter { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { if !arena.next_tuple(self.input, plan_arena)? { diff --git a/src/execution/dql/function_scan.rs b/src/execution/dql/function_scan.rs index 2e43d6f5..eb54d99c 100644 --- a/src/execution/dql/function_scan.rs +++ b/src/execution/dql/function_scan.rs @@ -16,6 +16,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::expression::function::table::TableFunction; use crate::planner::operator::function_scan::FunctionScanOperator; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; @@ -39,7 +40,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for FunctionScan { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -52,7 +53,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for FunctionScan { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.iter.is_none() { let TableFunction { args, catalog } = &self.table_function; diff --git a/src/execution/dql/index_scan.rs b/src/execution/dql/index_scan.rs index d2714acc..488b6ed5 100644 --- a/src/execution/dql/index_scan.rs +++ b/src/execution/dql/index_scan.rs @@ -16,6 +16,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::expression::range_detacher::Range; use crate::planner::operator::table_scan::TableScanOperator; +use crate::planner::MetaArena; use crate::storage::{IndexIter, IndexRanges, Iter, Transaction}; use crate::types::index::{IndexLookup, IndexMetaRef, RuntimeIndexProbe}; use crate::types::serialize::TupleValueSerializableImpl; @@ -64,7 +65,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for IndexScan<'a, T> { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -77,7 +78,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for IndexScan<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.iter.is_none() { let Some(TableScanOperator { diff --git a/src/execution/dql/join/hash/full_join.rs b/src/execution/dql/join/hash/full_join.rs index 6916c9ce..64a96948 100644 --- a/src/execution/dql/join/hash/full_join.rs +++ b/src/execution/dql/join/hash/full_join.rs @@ -18,7 +18,8 @@ use crate::execution::dql::join::hash::{ }; use crate::execution::dql::join::hash_join::BuildState; use crate::execution::dql::join::RowBitmap; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; @@ -34,7 +35,7 @@ impl JoinProbeState for FullJoinState { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { if probe_state.emitted_unmatched { @@ -102,7 +103,7 @@ impl JoinProbeState for FullJoinState { &mut self, left_drop_state: &mut LeftDropState, _filter_expr: Option<&ExprRef>, - _plan_arena: &PlanArena<'_>, + _plan_arena: &(dyn MetaArena + '_), ) -> 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 74ea4574..95066e84 100644 --- a/src/execution/dql/join/hash/inner_join.rs +++ b/src/execution/dql/join/hash/inner_join.rs @@ -15,7 +15,8 @@ use crate::errors::DatabaseError; use crate::execution::dql::join::hash::{filter, JoinProbeState, ProbeState}; use crate::execution::dql::join::hash_join::BuildState; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::{SplitTupleRef, Tuple}; pub(crate) struct InnerJoinState; @@ -26,7 +27,7 @@ impl JoinProbeState for InnerJoinState { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { probe_state.finished = true; diff --git a/src/execution/dql/join/hash/left_join.rs b/src/execution/dql/join/hash/left_join.rs index 103d536e..dd99951d 100644 --- a/src/execution/dql/join/hash/left_join.rs +++ b/src/execution/dql/join/hash/left_join.rs @@ -18,7 +18,8 @@ use crate::execution::dql::join::hash::{ }; use crate::execution::dql::join::hash_join::BuildState; use crate::execution::dql::join::RowBitmap; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; @@ -34,7 +35,7 @@ impl JoinProbeState for LeftJoinState { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { probe_state.finished = true; @@ -82,7 +83,7 @@ impl JoinProbeState for LeftJoinState { &mut self, left_drop_state: &mut LeftDropState, _filter_expr: Option<&ExprRef>, - _plan_arena: &PlanArena<'_>, + _plan_arena: &(dyn MetaArena + '_), ) -> 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 2ead9d32..a93478f1 100644 --- a/src/execution/dql/join/hash/mod.rs +++ b/src/execution/dql/join/hash/mod.rs @@ -24,7 +24,8 @@ 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::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::{Tuple, TupleLike}; use crate::types::value::DataValue; use std::collections::hash_map::IntoIter as HashMapIntoIter; @@ -55,14 +56,14 @@ pub(crate) trait JoinProbeState { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError>; fn left_drop_next( &mut self, _left_drop_state: &mut LeftDropState, _filter_expr: Option<&ExprRef>, - _plan_arena: &PlanArena<'_>, + _plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { Ok(None) } @@ -81,7 +82,7 @@ impl JoinProbeState for JoinProbeStateImpl { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { match self { JoinProbeStateImpl::Inner(state) => { @@ -103,7 +104,7 @@ impl JoinProbeState for JoinProbeStateImpl { &mut self, left_drop_state: &mut LeftDropState, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { match self { JoinProbeStateImpl::Inner(state) => { @@ -125,9 +126,9 @@ impl JoinProbeState for JoinProbeStateImpl { pub(crate) fn filter( values: &T, filter_expr: &ExprRef, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result { - match &plan_arena + match &*plan_arena .expression(*filter_expr) .eval(plan_arena, Some(values as &dyn TupleLike))? { diff --git a/src/execution/dql/join/hash/right_join.rs b/src/execution/dql/join/hash/right_join.rs index b2548c5d..0f40c4cc 100644 --- a/src/execution/dql/join/hash/right_join.rs +++ b/src/execution/dql/join/hash/right_join.rs @@ -16,7 +16,8 @@ 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::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::{SplitTupleRef, Tuple}; pub(crate) struct RightJoinState { @@ -29,7 +30,7 @@ impl JoinProbeState for RightJoinState { probe_state: &mut ProbeState, build_state: Option<&mut BuildState>, filter_expr: Option<&ExprRef>, - plan_arena: &PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { if probe_state.is_keys_has_null { if probe_state.emitted_unmatched { diff --git a/src/execution/dql/join/hash_join.rs b/src/execution/dql/join/hash_join.rs index 15e35101..d44d1d37 100644 --- a/src/execution/dql/join/hash_join.rs +++ b/src/execution/dql/join/hash_join.rs @@ -26,13 +26,14 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::join::{JoinCondition, JoinOperator, JoinType}; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use bumpalo::Bump; use std::collections::HashMap; -use std::mem::{self, transmute}; +use std::mem::transmute; pub struct HashJoin { state: HashJoinState, @@ -130,11 +131,16 @@ impl HashJoin { on_keys: &[ExprRef], tuple: &Tuple, build_buf: &mut BumpVec<'_, DataValue>, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { build_buf.clear(); for expr in on_keys { - build_buf.push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); + build_buf.push( + plan_arena + .expression(*expr) + .eval(plan_arena, Some(tuple))? + .into_owned(), + ); } Ok(()) } @@ -142,7 +148,7 @@ impl HashJoin { fn initialize_build<'a, T: Transaction + 'a>( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if !matches!(self.state, HashJoinState::Build) { return Ok(()); @@ -157,7 +163,7 @@ impl HashJoin { let mut build_count = 0usize; while arena.next_tuple(self.left_input, plan_arena)? { - let tuple = mem::take(arena.result_tuple_mut()); + let tuple = arena.materialize_tuple(); Self::eval_keys(&self.on_left_keys, &tuple, &mut build_buf, plan_arena)?; match build_map.get_mut(&build_buf) { @@ -229,7 +235,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for HashJoin { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -260,7 +266,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashJoin { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if let Some(err) = self.init_error.take() { return Err(err); @@ -283,7 +289,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for HashJoin { if !arena.next_tuple(self.right_input, plan_arena)? { break true; } - let tuple = mem::take(arena.result_tuple_mut()); + let tuple = arena.materialize_tuple(); Self::eval_keys( &self.on_right_keys, &tuple, @@ -385,6 +391,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::test::PlanArenaTestExt; use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::Storage; @@ -429,8 +436,8 @@ mod test { let on_keys = vec![(left_key, right_key)]; let values_t1 = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[ vec![ DataValue::Int32(0), DataValue::Int32(2), @@ -446,15 +453,16 @@ mod test { DataValue::Int32(5), DataValue::Int32(7), ], - ], - schema_ref: t1_columns, - }), + ]), + 3, + t1_columns, + )), Childrens::None, ); let values_t2 = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[ vec![ DataValue::Int32(0), DataValue::Int32(2), @@ -475,9 +483,10 @@ mod test { DataValue::Int32(1), DataValue::Int32(1), ], - ], - schema_ref: t2_columns, - }), + ]), + 4, + t2_columns, + )), Childrens::None, ); @@ -705,20 +714,22 @@ mod test { }); let left = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![DataValue::Int32(2), DataValue::Int32(0)], vec![DataValue::Int32(2), DataValue::Int32(5)], - ], - schema_ref: left_columns, - }), + ]), + 2, + left_columns, + )), Childrens::None, ); let right = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![vec![DataValue::Int32(2)]], - schema_ref: right_columns, - }), + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[vec![DataValue::Int32(2)]]), + 1, + right_columns, + )), Childrens::None, ); diff --git a/src/execution/dql/join/nested_loop_join.rs b/src/execution/dql/join/nested_loop_join.rs index 0847a125..a0aa111a 100644 --- a/src/execution/dql/join/nested_loop_join.rs +++ b/src/execution/dql/join/nested_loop_join.rs @@ -15,7 +15,7 @@ //! Defines the nested loop join executor, it supports [`JoinType::Inner`], [`JoinType::LeftOuter`], //! [`JoinType::RightOuter`], [`JoinType::Cross`], [`JoinType::Full`]. -use std::mem; +use crate::planner::MetaArena; use crate::errors::DatabaseError; use crate::execution::dql::join::RowBitmap; @@ -24,7 +24,7 @@ use crate::execution::{ }; use crate::iter_ext::Itertools; use crate::planner::operator::join::{JoinCondition, JoinOperator, JoinType}; -use crate::planner::{ExprRef, LogicalPlan, PlanArena}; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; use crate::types::tuple::{SplitTupleRef, Tuple}; use crate::types::value::DataValue; @@ -45,7 +45,7 @@ impl EqualCondition { &self, left_tuple: &Tuple, right_tuple: &Tuple, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result { if self.on_left_keys.is_empty() { return Ok(true); @@ -81,6 +81,7 @@ pub struct NestedLoopJoin { filter: Option, eq_cond: EqualCondition, left_input: ExecId, + right_pos: ExecId, state: NestedLoopJoinState, } @@ -142,6 +143,7 @@ impl From<(JoinOperator, LogicalPlan, LogicalPlan)> for NestedLoopJoin { filter, eq_cond, left_input: 0, + right_pos: 0, state: NestedLoopJoinState::PullLeft { right_bitmap: None }, } } @@ -153,7 +155,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for NestedLoopJoin { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -169,6 +171,14 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for NestedLoopJoin { cache, transaction, ); + executor.right_pos = arena.nodes.position(); + build_read( + arena, + plan_arena, + executor.right_input_plan.clone(), + cache, + transaction, + ); arena.push(ExecNode::NestedLoopJoin(executor)) } } @@ -177,7 +187,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let mut state = std::mem::replace(&mut self.state, NestedLoopJoinState::End); @@ -197,7 +207,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { arena.finish(); return Ok(()); } - let left_tuple = mem::take(arena.result_tuple_mut()); + let left_tuple = arena.materialize_tuple(); state = NestedLoopJoinState::ScanRight { active_left: ActiveLeftState { @@ -215,7 +225,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { mut right_bitmap, } => { while arena.next_tuple(active_left.right_input, plan_arena)? { - let right_tuple = mem::take(arena.result_tuple_mut()); + let right_tuple = arena.materialize_tuple(); let idx = active_left.right_index; active_left.right_index += 1; @@ -253,8 +263,8 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { }; let value = plan_arena .expression(*filter) - .eval(plan_arena, Some(values))?; - match &value { + .eval(plan_arena, Some(&values))?; + match &*value { DataValue::Boolean(true) => { let tuple = match self.ty { JoinType::RightOuter => Self::emit_tuple( @@ -350,7 +360,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for NestedLoopJoin { mut right_emit_index, } => { while arena.next_tuple(right_input, plan_arena)? { - let mut right_tuple = mem::take(arena.result_tuple_mut()); + let mut right_tuple = arena.materialize_tuple(); let idx = right_emit_index; right_emit_index += 1; @@ -390,11 +400,12 @@ impl NestedLoopJoin { fn build_right_input<'a, T: Transaction + 'a>( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> ExecId { let cache = arena.context(); let transaction = arena.transaction(); - // Fixme: Executor reset + // The same right-hand plan rebuilds the same slots, including nested joins. + arena.nodes.seek(self.right_pos); build_read( arena, plan_arena, @@ -463,6 +474,7 @@ mod test { use crate::optimizer::rule::normalization::NormalizationRuleImpl; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::Childrens; use crate::storage::rocksdb::RocksStorage; use crate::storage::Storage; @@ -533,8 +545,8 @@ mod test { }; let values_t1 = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[ vec![ DataValue::Int32(0), DataValue::Int32(2), @@ -555,15 +567,16 @@ mod test { DataValue::Int32(5), DataValue::Int32(7), ], - ], - schema_ref: t1_columns, - }), + ]), + 4, + t1_columns, + )), Childrens::None, ); let values_t2 = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[ vec![ DataValue::Int32(0), DataValue::Int32(2), @@ -584,9 +597,10 @@ mod test { DataValue::Int32(1), DataValue::Int32(1), ], - ], - schema_ref: t2_columns, - }), + ]), + 4, + t2_columns, + )), Childrens::None, ); @@ -636,6 +650,50 @@ mod test { assert!(expected.is_empty()); } + #[test] + fn nested_right_subtrees_reuse_slots() -> Result<(), DatabaseError> { + let storage = crate::storage::memory::MemoryStorage::new(); + let transaction = storage.transaction()?; + let meta_cache = crate::storage::StatisticsMetaCache::default(); + let view_cache = crate::storage::ViewCache::default(); + let table_cache = crate::storage::TableCache::default(); + let table_arena = crate::planner::TableArenaCell::default(); + let mut plan_arena = crate::planner::PlanArena::new(&table_arena); + let (_, left, right, _) = build_join_values(&mut plan_arena, false); + let cross = |left, right| { + LogicalPlan::new( + Operator::Join(JoinOperator { + on: JoinCondition::None, + join_type: JoinType::Cross, + force_nested_loop: true, + }), + Childrens::Twins { + left: Box::new(left), + right: Box::new(right), + }, + ) + }; + let plan = cross(left.clone(), cross(left, right)); + let context = crate::execution::empty_context(&table_cache, &view_cache, &meta_cache); + let mut arena = ExecArena::new(); + arena.init_context(context, &transaction); + let root = build_read(&mut arena, &mut plan_arena, plan, context, &transaction); + let count = arena.nodes.items.len(); + let address = arena.nodes.items.as_ptr(); + assert_eq!(count, 5); + let mut rows = 0; + while arena.next_tuple(root, &mut plan_arena)? { + rows += 1; + assert_eq!(arena.result_tuple().values.len(), 9); + assert_eq!(arena.nodes.items.len(), count); + assert_eq!(arena.nodes.items.as_ptr(), address); + } + assert_eq!(rows, 64); + assert_eq!(arena.nodes.items.len(), count); + assert!(!arena.next_tuple(root, &mut plan_arena)?); + Ok(()) + } + #[test] fn test_nested_inner_join() -> Result<(), DatabaseError> { let temp_dir = TempDir::new().expect("unable to create temporary working directory"); @@ -1160,20 +1218,22 @@ mod test { }); let left = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![ + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[ vec![DataValue::Int32(2), DataValue::Int32(0)], vec![DataValue::Int32(2), DataValue::Int32(5)], - ], - schema_ref: left_columns, - }), + ]), + 2, + left_columns, + )), Childrens::None, ); let right = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![vec![DataValue::Int32(2)]], - schema_ref: right_columns, - }), + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[vec![DataValue::Int32(2)]]), + 1, + right_columns, + )), Childrens::None, ); diff --git a/src/execution/dql/limit.rs b/src/execution/dql/limit.rs index 20bdf39d..5ab6a1b8 100644 --- a/src/execution/dql/limit.rs +++ b/src/execution/dql/limit.rs @@ -18,6 +18,7 @@ use crate::execution::{ }; use crate::planner::operator::limit::LimitOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; pub struct Limit { offset: Option, @@ -33,7 +34,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Limit { fn into_executor( (LimitOperator { offset, limit }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -52,7 +53,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Limit { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let offset = self.offset.unwrap_or(0); let limit = self.limit.unwrap_or(usize::MAX); diff --git a/src/execution/dql/mark_apply.rs b/src/execution/dql/mark_apply.rs index 5f1a812f..3670c9a5 100644 --- a/src/execution/dql/mark_apply.rs +++ b/src/execution/dql/mark_apply.rs @@ -18,6 +18,7 @@ use crate::execution::{ }; use crate::planner::operator::mark_apply::{MarkApplyKind, MarkApplyOperator, MarkApplyQuantifier}; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::index::RuntimeIndexProbe; use crate::types::tuple::{SplitTupleRef, Tuple}; @@ -31,39 +32,43 @@ enum QuantifiedPredicateOutcome { Skip, } -pub struct MarkApply<'a, T: Transaction + 'a> { +pub struct MarkApply { op: MarkApplyOperator, right_input_plan: LogicalPlan, left_input: ExecId, + right_pos: ExecId, // Retain a streaming inner input across next_tuple calls, not its result rows. - join_input: Option<(Box>, ExecId, Tuple)>, + join_input: Option<(ExecId, Tuple)>, } -impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for MarkApply<'a, T> { +impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for MarkApply { type Input = (MarkApplyOperator, LogicalPlan, LogicalPlan); fn into_executor( (op, left_input, right_input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { let left_input = build_read(arena, plan_arena, left_input, cache, transaction); + let right_pos = arena.nodes.position(); + build_read(arena, plan_arena, right_input.clone(), cache, transaction); arena.push(ExecNode::MarkApply(Self { op, right_input_plan: right_input, left_input, + right_pos, join_input: None, })) } } -impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for MarkApply<'a, T> { +impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for MarkApply { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if matches!(self.op.kind, MarkApplyKind::InnerJoin) { return self.next_join_tuple(arena, plan_arena); @@ -73,7 +78,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for MarkApply<'a, T> { return Ok(()); } - let mut left_tuple = arena.result_tuple().clone(); + let mut left_tuple = arena.materialize_tuple(); let marker = self.mark_value(arena, plan_arena, &left_tuple)?; left_tuple.values.push(marker); @@ -82,16 +87,16 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for MarkApply<'a, T> { } } -impl<'a, T: Transaction + 'a> MarkApply<'a, T> { - fn next_join_tuple( +impl MarkApply { + fn next_join_tuple<'a, T: Transaction + 'a>( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { - if let Some((inner, root, left)) = &mut self.join_input { - while inner.next_tuple(*root, plan_arena)? { - let right = inner.result_tuple(); + if let Some((root, left)) = &mut self.join_input { + while arena.next_tuple(*root, plan_arena)? { + let right = arena.result_tuple(); if Self::predicates_matched(self.op.predicates(), left, right, plan_arena)? { let mut output = left.clone(); output.pk = output.pk.or_else(|| right.pk.clone()); @@ -106,17 +111,10 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { arena.finish(); return Ok(()); } - let left: Tuple = arena.result_tuple().clone(); + let left: Tuple = arena.materialize_tuple(); let value = self.parameterized_probe_value(&left, plan_arena)?; - let mut inner = self - .join_input - .take() - .map(|(inner, _, _)| inner) - .unwrap_or_else(|| Box::new(ExecArena::new())); - inner.reset_for_rebuild(); - inner.init_context(arena.context(), arena.transaction()); - let root = self.build_right_input(&mut inner, plan_arena, value); - self.join_input = Some((inner, root, left)); + let root = self.build_right_input(arena, plan_arena, value); + self.join_input = Some((root, left)); } } @@ -139,12 +137,13 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { } } - fn build_right_input( + fn build_right_input<'a, T: Transaction + 'a>( &self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), param_value: Option, ) -> ExecId { + arena.nodes.seek(self.right_pos); if let Some(probe) = self.runtime_probe_for(param_value) { arena.push_runtime_probe(probe); } @@ -157,37 +156,10 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { ) } - fn with_right_input( - &self, - arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, - param_value: Option, - f: impl FnOnce( - &mut ExecArena<'a, T>, - &mut crate::planner::PlanArena<'a>, - ExecId, - ) -> Result, - ) -> Result { - let depth_before = arena.runtime_probe_depth(); - let right_input = self.build_right_input(arena, plan_arena, param_value); - let result = f(arena, plan_arena, right_input); - - let depth_after = arena.runtime_probe_depth(); - debug_assert!( - depth_after == depth_before || depth_after == depth_before + 1, - "parameterized right input should consume at most one runtime probe" - ); - if depth_after > depth_before { - let _ = arena.pop_runtime_probe(); - } - - result - } - fn parameterized_probe_value( &self, left_tuple: &Tuple, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { self.op .parameterized_probe() @@ -195,119 +167,96 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { plan_arena .expression(*probe) .eval(plan_arena, Some(left_tuple)) + .map(|value| value.into_owned()) }) .transpose() } - fn mark_value( + fn mark_value<'a, T: Transaction + 'a>( &self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), left_tuple: &Tuple, ) -> Result { + let probe = self.parameterized_probe_value(left_tuple, plan_arena)?; match self.op.kind { MarkApplyKind::InnerJoin => unreachable!("inner join streams tuples"), - MarkApplyKind::Exists => self.with_right_input( - arena, - plan_arena, - 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::predicates_matched( - self.op.predicates(), - left_tuple, - right_tuple, - plan_arena, - )? { - return Ok(DataValue::Boolean(true)); - } + MarkApplyKind::Exists => { + let right_input = self.build_right_input(arena, plan_arena, probe); + while arena.next_tuple(right_input, plan_arena)? { + let right_tuple = arena.result_tuple(); + if Self::predicates_matched( + self.op.predicates(), + left_tuple, + right_tuple, + plan_arena, + )? { + return Ok(DataValue::Boolean(true)); } - - Ok(DataValue::Boolean(false)) - }, - ), + } + Ok(DataValue::Boolean(false)) + } MarkApplyKind::Quantified(MarkApplyQuantifier::Any) => { if let Some(probe_value) = self.parameterized_probe_value(left_tuple, plan_arena)? { if !probe_value.is_null() { - if self.with_right_input( - arena, - plan_arena, - Some(probe_value), - |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, - plan_arena, - )? == QuantifiedPredicateOutcome::True - { - return Ok(true); - } - } - - Ok(false) - }, - )? { - return Ok(DataValue::Boolean(true)); + let right_input = + self.build_right_input(arena, plan_arena, Some(probe_value)); + while arena.next_tuple(right_input, plan_arena)? { + let right_tuple = arena.result_tuple(); + if self.quantified_predicate_outcome( + left_tuple, + right_tuple, + plan_arena, + )? == QuantifiedPredicateOutcome::True + { + return Ok(DataValue::Boolean(true)); + } } - if self.with_right_input( - arena, - plan_arena, - Some(DataValue::Null), - |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, - plan_arena, - )? == QuantifiedPredicateOutcome::Null - { - return Ok(true); - } - } - - Ok(false) - }, - )? { - return Ok(DataValue::Null); + let right_input = + self.build_right_input(arena, plan_arena, Some(DataValue::Null)); + while arena.next_tuple(right_input, plan_arena)? { + let right_tuple = arena.result_tuple(); + if self.quantified_predicate_outcome( + left_tuple, + right_tuple, + plan_arena, + )? == QuantifiedPredicateOutcome::Null + { + return Ok(DataValue::Null); + } } return Ok(DataValue::Boolean(false)); } } - self.with_right_input(arena, plan_arena, None, |arena, plan_arena, right_input| { - self.scan_quantified_right_input( - arena, - plan_arena, - right_input, - MarkApplyQuantifier::Any, - left_tuple, - ) - }) + let right_input = self.build_right_input(arena, plan_arena, None); + self.scan_quantified_right_input( + arena, + plan_arena, + right_input, + MarkApplyQuantifier::Any, + left_tuple, + ) } MarkApplyKind::Quantified(MarkApplyQuantifier::All) => { - self.with_right_input(arena, plan_arena, None, |arena, plan_arena, right_input| { - self.scan_quantified_right_input( - arena, - plan_arena, - right_input, - MarkApplyQuantifier::All, - left_tuple, - ) - }) + let right_input = self.build_right_input(arena, plan_arena, None); + self.scan_quantified_right_input( + arena, + plan_arena, + right_input, + MarkApplyQuantifier::All, + left_tuple, + ) } } } - fn scan_quantified_right_input( + fn scan_quantified_right_input<'a, T: Transaction + 'a>( &self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), right_input: ExecId, quantifier: MarkApplyQuantifier, left_tuple: &Tuple, @@ -346,14 +295,14 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { predicates: &[crate::planner::ExprRef], left_tuple: &Tuple, right_tuple: &Tuple, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result { let values = SplitTupleRef::new(left_tuple, right_tuple); for predicate in predicates { - match plan_arena + match *plan_arena .expression(*predicate) - .eval(plan_arena, Some(values))? + .eval(plan_arena, Some(&values))? { DataValue::Boolean(true) => {} DataValue::Boolean(false) | DataValue::Null => return Ok(false), @@ -368,9 +317,9 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { &self, left_tuple: &Tuple, right_tuple: &Tuple, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result { - match self.eval_predicates(left_tuple, right_tuple, plan_arena)? { + 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), @@ -383,7 +332,7 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { &self, left_tuple: &Tuple, right_tuple: &Tuple, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { let values = SplitTupleRef::new(left_tuple, right_tuple); // probe_predicate is in predicate, always first @@ -394,9 +343,9 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { .ok_or(DatabaseError::InvalidType)?; for predicate in correlated_predicates { - match plan_arena + match *plan_arena .expression(*predicate) - .eval(plan_arena, Some(values))? + .eval(plan_arena, Some(&values))? { DataValue::Boolean(true) => {} DataValue::Boolean(false) | DataValue::Null => return Ok(None), @@ -407,7 +356,8 @@ impl<'a, T: Transaction + 'a> MarkApply<'a, T> { Ok(Some( plan_arena .expression(*probe_predicate) - .eval(plan_arena, Some(values))?, + .eval(plan_arena, Some(&values))? + .into_owned(), )) } } @@ -420,6 +370,7 @@ mod tests { use crate::expression::{BinaryOperator, ScalarExpression}; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::{Childrens, ExprRef, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::{StatisticsMetaCache, Storage, TableCache, ViewCache}; @@ -447,7 +398,11 @@ mod tests { .collect(); LogicalPlan::new( - Operator::Values(ValuesOperator { rows, schema_ref }), + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&rows), + rows.len(), + schema_ref, + )), Childrens::None, ) } @@ -512,7 +467,7 @@ mod tests { } #[test] - fn inner_join_apply_emits_all_matches_and_reuses_inner_arena() -> Result<(), DatabaseError> { + fn inner_join_apply_emits_all_matches_and_reuses_right_slots() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); let mut plan_arena = crate::planner::PlanArena::new(&table_arena); let mut left = build_values( @@ -547,43 +502,38 @@ mod tests { plan_arena.alloc_expression(ScalarExpression::column_expr(right_schema[1], 2)); let probe = plan_arena.alloc_expression(ScalarExpression::column_expr(left_column, 0)); let mut op = MarkApplyOperator::new_inner_join(vec![equality, residual], probe); - // Values supplies the inner rows directly; there is no IndexScan to consume a probe. + // Values has no parameterized IndexScan to consume a runtime probe. op.set_parameterized_probe(None); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; let cache = crate::execution::empty_context(&table_cache, &view_cache, &meta_cache); let mut arena = ExecArena::new(); arena.init_context(cache, &transaction); - let left_input = build_read(&mut arena, &mut plan_arena, left, cache, &transaction); - let mut exec = MarkApply { - op, - right_input_plan: right, - left_input, - join_input: None, - }; - let mut inner_address = None; + let root = >::into_executor( + (op, left, right), + &mut arena, + &mut plan_arena, + cache, + &transaction, + ); + let address = arena.nodes.items.as_ptr(); + let count = arena.nodes.items.len(); + assert_eq!(count, 3); for _ in 0..4 { - exec.next_tuple(&mut arena, &mut plan_arena)?; + assert!(arena.next_tuple(root, &mut plan_arena)?); assert_eq!( - arena.result_tuple().values, + arena.materialize_tuple().values, vec![ DataValue::Int32(2), DataValue::Int32(2), DataValue::Boolean(true) ] ); - let (inner, _, _) = exec.join_input.as_ref().expect("active inner scan"); - let address = &**inner as *const _; - assert_eq!(*inner_address.get_or_insert(address), address); - assert_eq!( - inner.nodes.len(), - 1, - "inner executors must not accumulate per outer row" - ); - assert_eq!(inner.runtime_probe_depth(), 0); + assert_eq!(arena.nodes.items.as_ptr(), address); + assert_eq!(arena.nodes.items.len(), count); + assert_eq!(arena.runtime_probe_depth(), 0); } - exec.next_tuple(&mut arena, &mut plan_arena)?; - assert!(exec.join_input.is_none(), "all outer rows exhausted"); + assert!(!arena.next_tuple(root, &mut plan_arena)?); Ok(()) } @@ -609,7 +559,7 @@ mod tests { op.set_parameterized_probe(None); let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let mut executor = execute_input::<_, MarkApply<_>>( + let mut executor = execute_input::<_, MarkApply>( (op, left, right), crate::execution::empty_context(&table_cache, &view_cache, &meta_cache), plan_arena, @@ -644,7 +594,10 @@ mod tests { let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let tuples = try_collect(execute_input::<_, MarkApply<_>>( + let context = crate::execution::empty_context(&table_cache, &view_cache, &meta_cache); + let mut arena = ExecArena::new(); + arena.init_context(context, &transaction); + let root = >::into_executor( ( MarkApplyOperator::new_exists( build_marker_column(&mut plan_arena), @@ -653,10 +606,21 @@ mod tests { left, right, ), - crate::execution::empty_context(&table_cache, &view_cache, &meta_cache), - plan_arena, + &mut arena, + &mut plan_arena, + context, &transaction, - ))?; + ); + let count = arena.nodes.items.len(); + let address = arena.nodes.items.as_ptr(); + assert_eq!(count, 3); + let mut tuples = Vec::new(); + while arena.next_tuple(root, &mut plan_arena)? { + tuples.push(arena.materialize_tuple()); + assert_eq!(arena.nodes.items.len(), count); + assert_eq!(arena.nodes.items.as_ptr(), address); + assert_eq!(arena.runtime_probe_depth(), 0); + } assert_eq!( tuples @@ -695,7 +659,7 @@ mod tests { let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let tuples = try_collect(execute_input::<_, MarkApply<_>>( + let tuples = try_collect(execute_input::<_, MarkApply>( ( MarkApplyOperator::new_exists( build_marker_column(&mut plan_arena), @@ -776,10 +740,11 @@ mod tests { &transaction, ); - let exec: MarkApply = MarkApply { + let exec: MarkApply = MarkApply { op, right_input_plan: right, left_input: 0, + right_pos: arena.nodes.position(), join_input: None, }; let left_tuple = Tuple::new(None, vec![DataValue::Int32(2), DataValue::Int32(1)]); @@ -828,10 +793,11 @@ mod tests { &transaction, ); - let exec: MarkApply = MarkApply { + let exec: MarkApply = MarkApply { op, right_input_plan: right, left_input: 0, + right_pos: arena.nodes.position(), join_input: None, }; let left_tuple = Tuple::new(None, vec![DataValue::Int32(2)]); @@ -880,10 +846,11 @@ mod tests { &transaction, ); - let exec: MarkApply = MarkApply { + let exec: MarkApply = MarkApply { op, right_input_plan: right, left_input: 0, + right_pos: arena.nodes.position(), join_input: None, }; let left_tuple = Tuple::new(None, vec![DataValue::Null]); @@ -924,7 +891,7 @@ mod tests { let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let tuples = try_collect(execute_input::<_, MarkApply<_>>( + let tuples = try_collect(execute_input::<_, MarkApply>( ( MarkApplyOperator::new_in(build_marker_column(&mut plan_arena), vec![predicate]), left, @@ -972,7 +939,7 @@ mod tests { let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let tuples = try_collect(execute_input::<_, MarkApply<_>>( + let tuples = try_collect(execute_input::<_, MarkApply>( ( MarkApplyOperator::new_in(build_marker_column(&mut plan_arena), vec![predicate]), left, @@ -1040,7 +1007,7 @@ mod tests { let (table_cache, view_cache, meta_cache, _temp_dir, storage) = build_test_storage()?; let transaction = storage.transaction()?; - let tuples = try_collect(execute_input::<_, MarkApply<_>>( + let tuples = try_collect(execute_input::<_, MarkApply>( ( MarkApplyOperator::new_in( build_marker_column(&mut plan_arena), diff --git a/src/execution/dql/projection.rs b/src/execution/dql/projection.rs index b30e7cc3..7817dd88 100644 --- a/src/execution/dql/projection.rs +++ b/src/execution/dql/projection.rs @@ -17,6 +17,7 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::project::ProjectOperator; +use crate::planner::MetaArena; use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::Transaction; @@ -31,7 +32,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Projection { fn into_executor( (ProjectOperator { exprs }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -44,22 +45,14 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Projection { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if !arena.next_tuple(self.input, plan_arena)? { arena.finish(); return Ok(()); } - arena.with_projection_tmp(|arena, projection_tmp| { - let tuple = arena.result_tuple(); - projection_tmp.reserve(self.exprs.len()); - for expr in self.exprs.iter() { - 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>(()) - })?; + arena.rewrite(&self.exprs, plan_arena, None)?; arena.resume(); Ok(()) } diff --git a/src/execution/dql/recursive_cte.rs b/src/execution/dql/recursive_cte.rs index a3fc5010..519075dd 100644 --- a/src/execution/dql/recursive_cte.rs +++ b/src/execution/dql/recursive_cte.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -19,7 +21,8 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::operator::recursive_cte::RecursiveScanOperator; -use crate::planner::{LogicalPlan, PlanArena}; +use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; use std::mem; @@ -198,7 +201,10 @@ impl<'a, T: Transaction + 'a> RecursiveCte<'a, T> { } } - fn start_recursive(&mut self, plan_arena: &mut PlanArena<'a>) -> Result { + fn start_recursive( + &mut self, + plan_arena: &mut (dyn MetaArena + 'a), + ) -> Result { let Some(input) = mem::take(&mut self.working).into_input()? else { return Ok(false); }; @@ -224,7 +230,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for RecursiveCte<'a, T> { fn into_executor( (anchor_plan, recursive_plan): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -243,13 +249,13 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for RecursiveCte<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { match self.phase { RecursivePhase::Anchor => { while arena.next_tuple(self.anchor_input, plan_arena)? { - self.next.push(mem::take(arena.result_tuple_mut()))?; + self.next.push(arena.materialize_tuple())?; } self.working = mem::take(&mut self.next).finish()?; self.phase = RecursivePhase::Output; @@ -270,8 +276,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for RecursiveCte<'a, T> { .recursive_arena .next_tuple(self.recursive_root, plan_arena)? { - self.next - .push(mem::take(self.recursive_arena.result_tuple_mut()))?; + self.next.push(self.recursive_arena.materialize_tuple())?; } self.recursive_arena.reset_for_rebuild(); self.working = mem::take(&mut self.next).finish()?; @@ -292,7 +297,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for RecursiveScan { fn into_executor( _input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _cache: ExecutionContext<'_>, _transaction: &T, ) -> ExecId { @@ -305,7 +310,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for RecursiveScan { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { match self.input.next().transpose()? { Some(tuple) => arena.produce_tuple(tuple), @@ -326,6 +331,7 @@ mod tests { use crate::planner::operator::recursive_cte::RecursiveScanOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::Childrens; use crate::storage::rocksdb::RocksStorage; use crate::storage::{StatisticsMetaCache, Storage, TableCache, ViewCache}; @@ -402,10 +408,11 @@ mod tests { )); let schema_ref = vec![column]; let anchor = LogicalPlan::new( - Operator::Values(ValuesOperator { - rows: vec![vec![DataValue::Int32(1)]], - schema_ref: schema_ref.clone(), - }), + Operator::Values(ValuesOperator::new( + plan_arena.alloc_expression_rows(&[vec![DataValue::Int32(1)]]), + 1, + schema_ref.clone(), + )), Childrens::None, ); let scan = LogicalPlan::new( diff --git a/src/execution/dql/scalar_apply.rs b/src/execution/dql/scalar_apply.rs index 3f80ef44..7589fe6a 100644 --- a/src/execution/dql/scalar_apply.rs +++ b/src/execution/dql/scalar_apply.rs @@ -18,9 +18,9 @@ use crate::execution::{ }; use crate::planner::operator::scalar_apply::ScalarApplyOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; -use std::mem; pub struct ScalarApply { left_input: ExecId, @@ -34,7 +34,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for ScalarApply { fn into_executor( (_, left_input, right_input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -52,7 +52,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ScalarApply { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { Self::load_right_once(&mut self.cached_right, self.right_input, arena, plan_arena)?; @@ -78,7 +78,7 @@ impl ScalarApply { cached_right: &mut Option, right_input: ExecId, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if cached_right.is_none() { if !arena.next_tuple(right_input, plan_arena)? { @@ -86,7 +86,7 @@ impl ScalarApply { "scalar apply right input returned no rows".to_string(), )); } - *cached_right = Some(mem::take(arena.result_tuple_mut())); + *cached_right = Some(arena.materialize_tuple()); } Ok(()) @@ -101,6 +101,7 @@ mod tests { use crate::planner::operator::scalar_subquery::ScalarSubqueryOperator; use crate::planner::operator::values::ValuesOperator; use crate::planner::operator::Operator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::{Childrens, LogicalPlan}; use crate::storage::rocksdb::RocksStorage; use crate::storage::{StatisticsMetaCache, Storage, TableCache, ViewCache}; @@ -117,7 +118,11 @@ mod tests { let schema_ref = vec![arena.alloc_column(ColumnCatalog::new(name.to_string(), true, desc))]; LogicalPlan::new( - Operator::Values(ValuesOperator { rows, schema_ref }), + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&rows), + rows.len(), + schema_ref, + )), Childrens::None, ) } diff --git a/src/execution/dql/scalar_subquery.rs b/src/execution/dql/scalar_subquery.rs index 71253e11..d53fdfc4 100644 --- a/src/execution/dql/scalar_subquery.rs +++ b/src/execution/dql/scalar_subquery.rs @@ -18,6 +18,7 @@ use crate::execution::{ }; use crate::planner::operator::scalar_subquery::ScalarSubqueryOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -33,7 +34,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for ScalarSubquery { fn into_executor( (_, mut input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -51,7 +52,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ScalarSubquery { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.returned { arena.finish(); diff --git a/src/execution/dql/seq_scan.rs b/src/execution/dql/seq_scan.rs index 98b5513c..91df1a59 100644 --- a/src/execution/dql/seq_scan.rs +++ b/src/execution/dql/seq_scan.rs @@ -15,6 +15,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::planner::operator::table_scan::TableScanOperator; +use crate::planner::MetaArena; use crate::storage::{Iter, Transaction, TupleIter}; pub(crate) struct SeqScan<'a, T: Transaction + 'a> { @@ -37,7 +38,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for SeqScan<'a, T> { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -50,7 +51,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for SeqScan<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.iter.is_none() { let Some(TableScanOperator { diff --git a/src/execution/dql/set_membership.rs b/src/execution/dql/set_membership.rs index ecd041a0..fb2b6c55 100644 --- a/src/execution/dql/set_membership.rs +++ b/src/execution/dql/set_membership.rs @@ -18,10 +18,10 @@ use crate::execution::{ }; use crate::planner::operator::set_membership::SetMembershipKind; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; use std::collections::HashMap; -use std::mem; pub struct SetMembership { kind: SetMembershipKind, @@ -55,7 +55,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for SetMembership { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -82,13 +82,13 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for SetMembership { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if !self.built { while arena.next_tuple(self.right_input, plan_arena)? { *self .right_counts - .entry(mem::take(arena.result_tuple_mut())) + .entry(arena.materialize_tuple()) .or_insert(0) += 1; } self.built = true; diff --git a/src/execution/dql/show_table.rs b/src/execution/dql/show_table.rs index 758352a7..e523df1c 100644 --- a/src/execution/dql/show_table.rs +++ b/src/execution/dql/show_table.rs @@ -14,6 +14,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; +use crate::planner::MetaArena; use crate::storage::{TableIter, Transaction}; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; @@ -28,7 +29,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for ShowTables<'a, T> { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _: &mut crate::planner::PlanArena<'a>, + _: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -40,7 +41,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ShowTables<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.metas.is_none() { let mut state = arena.local_state(plan_arena); diff --git a/src/execution/dql/show_view.rs b/src/execution/dql/show_view.rs index e22ebd82..9f6ca4a6 100644 --- a/src/execution/dql/show_view.rs +++ b/src/execution/dql/show_view.rs @@ -14,6 +14,7 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; +use crate::planner::MetaArena; use crate::storage::{Transaction, ViewIter}; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; @@ -28,7 +29,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for ShowViews<'a, T> { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _: &mut crate::planner::PlanArena<'a>, + _: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -40,7 +41,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for ShowViews<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.metas.is_none() { let context = arena.context(); diff --git a/src/execution/dql/sort.rs b/src/execution/dql/sort.rs index cb7e28a8..807a550f 100644 --- a/src/execution/dql/sort.rs +++ b/src/execution/dql/sort.rs @@ -18,12 +18,13 @@ use crate::execution::{ }; use crate::planner::operator::sort::{SortField, SortOperator}; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use bumpalo::Bump; use std::cmp::Ordering; -use std::mem::{self, transmute, MaybeUninit}; +use std::mem::{transmute, MaybeUninit}; use std::ops::{Deref, DerefMut}; pub(crate) type BumpVec<'bump, T> = bumpalo::collections::Vec<'bump, T>; @@ -81,7 +82,7 @@ impl DerefMut for NullableVec<'_, T> { pub(crate) fn sort_tuples( sort_fields: &[SortField], tuples: &mut NullableVec<'_, (usize, Tuple)>, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { // Extract the results of calculating SortFields to avoid double calculation // of data during comparison. @@ -89,7 +90,12 @@ pub(crate) fn sort_tuples( for (x, SortField { expr, .. }) in sort_fields.iter().enumerate() { for (_, tuple) in tuples.iter() { - eval_values[x].push(plan_arena.expression(*expr).eval(plan_arena, Some(tuple))?); + eval_values[x].push( + plan_arena + .expression(*expr) + .eval(plan_arena, Some(tuple))? + .into_owned(), + ); } } @@ -155,7 +161,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Sort { fn into_executor( (SortOperator { sort_fields }, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -179,7 +185,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Sort { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { loop { if let Some((_, tuple)) = self.rows.pop() { @@ -188,7 +194,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Sort { } while arena.next_tuple(self.input, plan_arena)? { let offset = self.rows.len(); - self.rows.put((offset, mem::take(arena.result_tuple_mut()))); + self.rows.put((offset, arena.materialize_tuple())); } if self.rows.is_empty() { arena.finish(); diff --git a/src/execution/dql/top_k.rs b/src/execution/dql/top_k.rs index 6ceb11e0..23f1040f 100644 --- a/src/execution/dql/top_k.rs +++ b/src/execution/dql/top_k.rs @@ -20,13 +20,14 @@ use crate::execution::{ use crate::planner::operator::sort::SortField; use crate::planner::operator::top_k::TopKOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::table_codec::BumpBytes; use crate::storage::Transaction; use crate::types::tuple::Tuple; use bumpalo::Bump; use std::cmp::Ordering; use std::collections::{btree_set::IntoIter as BTreeSetIntoIter, BTreeSet}; -use std::mem::{self, transmute}; +use std::mem::transmute; #[derive(Eq, PartialEq, Debug)] struct CmpItem<'a> { @@ -53,7 +54,7 @@ fn top_sort<'a>( heap: &mut BTreeSet>, tuple: Tuple, keep_count: usize, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let mut full_key = BumpBytes::new_in(arena); for SortField { @@ -114,7 +115,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for TopK { input, ): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -134,7 +135,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for TopK { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.output.is_none() { let keep_count = self.offset.unwrap_or(0) + self.limit; @@ -146,7 +147,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for TopK { &self.arena, &self.sort_fields, &mut set, - mem::take(arena.result_tuple_mut()), + arena.materialize_tuple(), keep_count, plan_arena, )?; diff --git a/src/execution/dql/union.rs b/src/execution/dql/union.rs index 82484632..75fa2365 100644 --- a/src/execution/dql/union.rs +++ b/src/execution/dql/union.rs @@ -17,6 +17,7 @@ use crate::execution::{ build_read, ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor, }; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; pub struct Union { left_plan: LogicalPlan, @@ -44,7 +45,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Union { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -71,7 +72,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Union { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { if self.reading_left { if arena.next_tuple(self.left_input, plan_arena)? { diff --git a/src/execution/dql/values.rs b/src/execution/dql/values.rs index 8367754c..df7d81b1 100644 --- a/src/execution/dql/values.rs +++ b/src/execution/dql/values.rs @@ -15,20 +15,28 @@ use crate::errors::DatabaseError; use crate::execution::{ExecArena, ExecId, ExecNode, ExecutionContext, ExecutorNode, ReadExecutor}; use crate::planner::operator::values::ValuesOperator; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Schema; -use crate::types::value::DataValue; -use std::mem; pub struct Values { - rows: std::vec::IntoIter>, + rows: std::vec::IntoIter, + remaining_rows: usize, schema_ref: Schema, } impl From for Values { - fn from(ValuesOperator { rows, schema_ref }: ValuesOperator) -> Self { + fn from( + ValuesOperator { + rows, + row_count, + schema_ref, + }: ValuesOperator, + ) -> Self { Values { rows: rows.into_iter(), + remaining_rows: row_count, schema_ref, } } @@ -40,7 +48,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Values { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - _plan_arena: &mut crate::planner::PlanArena<'a>, + _plan_arena: &mut (dyn MetaArena + 'a), _: ExecutionContext<'_>, _: &T, ) -> ExecId { @@ -53,22 +61,29 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Values { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { - let Some(mut values) = self.rows.next() else { + if self.remaining_rows == 0 { arena.finish(); return Ok(()); - }; - - for (i, value) in values.iter_mut().enumerate() { - let ty = plan_arena.column(self.schema_ref[i]).datatype(); - - *value = mem::replace(value, DataValue::Null).cast(ty)?; } + self.remaining_rows -= 1; + let width = self.schema_ref.len(); let output = arena.result_tuple_mut(); output.pk = None; - output.values = values; + output.values.clear(); + for (i, expr) in self.rows.by_ref().take(width).enumerate() { + let ty = plan_arena.column(self.schema_ref[i]).datatype(); + output.values.push( + plan_arena + .expression(expr) + .eval(plan_arena, None)? + .into_owned() + .cast(ty)?, + ); + } + arena.resume(); Ok(()) } diff --git a/src/execution/dql/window.rs b/src/execution/dql/window.rs index a4cc24ae..e31db988 100644 --- a/src/execution/dql/window.rs +++ b/src/execution/dql/window.rs @@ -20,10 +20,10 @@ use crate::expression::window::WindowFunctionKind; use crate::planner::operator::sort::SortField; use crate::planner::operator::window::WindowOperator; use crate::planner::LogicalPlan; +use crate::planner::MetaArena; use crate::storage::Transaction; use crate::types::tuple::Tuple; use crate::types::value::DataValue; -use std::mem; mod function; @@ -82,7 +82,7 @@ impl<'a, T: Transaction + 'a> ReadExecutor<'a, T> for Window { fn into_executor( (operator, input): Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId { @@ -123,13 +123,14 @@ impl Window { fn update_keys( &mut self, tuple: &Tuple, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result, DatabaseError> { let mut boundary = (!self.state.started).then_some(Boundary::Partition); for (index, field) in self.sort_fields.iter().enumerate() { let value = plan_arena .expression(field.expr) - .eval(plan_arena, Some(tuple))?; + .eval(plan_arena, Some(tuple))? + .into_owned(); if self.state.started && self.state.sort_values[index] != value { if index < self.partition_by_len { boundary = Some(Boundary::Partition); @@ -152,10 +153,7 @@ impl Window { Ok(()) } - fn eval_functions( - &mut self, - plan_arena: &crate::planner::PlanArena<'_>, - ) -> Result<(), DatabaseError> { + fn eval_functions(&mut self, plan_arena: &(dyn MetaArena + '_)) -> Result<(), DatabaseError> { if self.state.buffered.is_empty() { return Ok(()); } @@ -183,7 +181,7 @@ impl Window { &mut self, tuple: Tuple, boundary: Option, - plan_arena: &crate::planner::PlanArena<'_>, + plan_arena: &(dyn MetaArena + '_), ) -> Result { let boundary = match boundary { Some(boundary) => Some(boundary), @@ -223,7 +221,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Window { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut crate::planner::PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { let mut output_ready = true; loop { @@ -241,7 +239,7 @@ impl<'a, T: Transaction + 'a> ExecutorNode<'a, T> for Window { let (tuple, boundary) = if let Some((tuple, boundary)) = self.state.pending.take() { (tuple, Some(boundary)) } else if arena.next_tuple(self.input, plan_arena)? { - (mem::take(arena.result_tuple_mut()), None) + (arena.materialize_tuple(), None) } else { self.eval_functions(plan_arena)?; self.input_exhausted = true; diff --git a/src/execution/dql/window/function.rs b/src/execution/dql/window/function.rs index e1b35550..53808f3d 100644 --- a/src/execution/dql/window/function.rs +++ b/src/execution/dql/window/function.rs @@ -16,7 +16,8 @@ use crate::errors::DatabaseError; use crate::execution::dql::aggregate::{create_accumulator, Accumulator}; use crate::expression::agg::AggKind; use crate::expression::window::WindowFunctionKind; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::Tuple; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -34,7 +35,7 @@ pub(super) trait WindowFunction { peer_start: usize, peer_index: usize, output_position: usize, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError>; } @@ -48,7 +49,7 @@ impl WindowFunction for RowNumber { _peer_start: usize, _peer_index: usize, output_position: usize, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for (row_index, row) in &mut rows[peer] { row.values[output_position] = DataValue::Int64((*row_index + 1) as i64); @@ -69,7 +70,7 @@ impl WindowFunction for Rank { peer_start: usize, peer_index: usize, output_position: usize, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let rank = if self.dense { peer_index + 1 @@ -103,13 +104,14 @@ impl WindowFunction for Aggregate { _peer_start: usize, _peer_index: usize, output_position: usize, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let Some(accumulator) = self.accumulator.as_mut() else { unreachable!() }; for (_, row) in &rows[peer.clone()] { - accumulator.update_value(&arena.expression(self.arg).eval(arena, Some(row))?)?; + accumulator + .update_value(arena.expression(self.arg).eval(arena, Some(row))?.as_ref())?; } accumulator.evaluate()?; let result = accumulator.result(); diff --git a/src/execution/mod.rs b/src/execution/mod.rs index 3bf60c08..dc339cf3 100644 --- a/src/execution/mod.rs +++ b/src/execution/mod.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -73,7 +75,8 @@ use crate::execution::dql::window::Window; use crate::expression::ScalarExpression; use crate::planner::operator::join::JoinCondition; use crate::planner::operator::{Operator, PhysicalOption, PlanImpl}; -use crate::planner::{LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{ExprRef, LogicalPlan}; use crate::storage::table_codec::TableCodec; use crate::storage::{StatisticsMetaCache, TableCache, Transaction, ViewCache}; use crate::types::index::RuntimeIndexProbe; @@ -141,6 +144,23 @@ pub(crate) struct ExecResult { pub(crate) status: Option, } +/// Resolves either an arena-backed expression or a direct scalar expression. +pub(crate) trait RewriteExpression { + fn expression<'a>(&'a self, arena: &'a dyn MetaArena) -> &'a ScalarExpression; +} + +impl RewriteExpression for ExprRef { + fn expression<'a>(&'a self, arena: &'a dyn MetaArena) -> &'a ScalarExpression { + arena.expression(*self) + } +} + +impl RewriteExpression for ScalarExpression { + fn expression<'a>(&'a self, _arena: &'a dyn MetaArena) -> &'a ScalarExpression { + self + } +} + pub struct Executor<'a, T: Transaction + 'a> { arena: ExecArena<'a, T>, root: ExecId, @@ -153,7 +173,7 @@ impl<'a, T: Transaction + 'a> Executor<'a, T> { pub(crate) fn next_tuple( &mut self, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result, DatabaseError> { if !self.arena.next_tuple(self.root, plan_arena)? { return Ok(None); @@ -195,7 +215,7 @@ pub(crate) enum ExecNode<'a, T: Transaction + 'a> { IndexScan(IndexScan<'a, T>), Insert(Insert), Limit(Limit), - MarkApply(MarkApply<'a, T>), + MarkApply(MarkApply), NestedLoopJoin(NestedLoopJoin), Projection(Projection), RecursiveCte(RecursiveCte<'a, T>), @@ -216,14 +236,13 @@ pub(crate) enum ExecNode<'a, T: Transaction + 'a> { Update(Update), Values(Values), Window(Window), - Empty, } pub(crate) trait ExecutorNode<'a, T: Transaction + 'a>: Sized { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError>; } @@ -231,7 +250,7 @@ impl<'a, T: Transaction + 'a> ExecNode<'a, T> { fn next_tuple( &mut self, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result<(), DatabaseError> { match self { ExecNode::AddColumn(exec) => { @@ -310,7 +329,7 @@ impl<'a, T: Transaction + 'a> ExecNode<'a, T> { >::next_tuple(exec, arena, plan_arena) } ExecNode::MarkApply(exec) => { - as ExecutorNode<'a, T>>::next_tuple(exec, arena, plan_arena) + >::next_tuple(exec, arena, plan_arena) } ExecNode::NestedLoopJoin(exec) => { >::next_tuple(exec, arena, plan_arena) @@ -372,16 +391,49 @@ impl<'a, T: Transaction + 'a> ExecNode<'a, T> { ExecNode::Window(exec) => { >::next_tuple(exec, arena, plan_arena) } - ExecNode::Empty => unreachable!("executor node re-entered while active"), } } } +struct ExecNodes<'a, T: Transaction + 'a> { + items: Vec>>, + pos: ExecId, + executing: usize, +} + +impl<'a, T: Transaction + 'a> ExecNodes<'a, T> { + fn position(&self) -> ExecId { + self.pos + } + + fn seek(&mut self, pos: ExecId) { + assert!(pos <= self.items.len(), "node position out of bounds"); + self.pos = pos; + } + + fn push(&mut self, node: ExecNode<'a, T>) -> ExecId { + let id = self.pos; + if id == self.items.len() { + assert_eq!(self.executing, 0, "cannot grow nodes during execution"); + self.items.push(std::cell::RefCell::new(node)); + } else { + *self.items[id].borrow_mut() = node; + } + self.pos += 1; + id + } + + fn clear(&mut self) { + assert_eq!(self.executing, 0, "cannot clear nodes during execution"); + self.items.clear(); + self.pos = 0; + } +} + pub(crate) struct ExecArena<'a, T: Transaction + 'a> { - nodes: Vec>, + nodes: ExecNodes<'a, T>, result: ExecResult, table_codec: TableCodec, - projection_tmp: Vec, context: Option>, transaction: *mut T, runtime_probe_stack: Vec, @@ -394,7 +446,7 @@ pub(crate) struct ExecArenaLocalState<'b, 'a, T: Transaction + 'a> { pub(crate) table_codec: &'b mut TableCodec, pub(crate) context: ExecutionContext<'a>, pub(crate) result: &'b mut ExecResult, - pub(crate) plan_arena: &'b PlanArena<'a>, + pub(crate) plan_arena: &'b (dyn MetaArena + 'a), ddl_apply: &'b mut Vec, } @@ -407,6 +459,18 @@ impl<'b, 'a, T: Transaction + 'a> ExecArenaLocalState<'b, 'a, T> { unsafe { (&mut *self.transaction, &mut *self.table_codec) } } + pub(crate) fn index_values_transaction_codec_mut( + &mut self, + ) -> (&[DataValue], &mut T, &mut TableCodec) { + unsafe { + ( + &self.result.tuple.values, + &mut *self.transaction, + &mut *self.table_codec, + ) + } + } + pub(crate) fn transaction_codec(&mut self) -> (&'a T, &mut TableCodec) { unsafe { (&*self.transaction, &mut *self.table_codec) } } @@ -427,10 +491,13 @@ impl<'b, 'a, T: Transaction + 'a> ExecArenaLocalState<'b, 'a, T> { impl<'a, T: Transaction + 'a> ExecArena<'a, T> { pub(crate) fn new() -> Self { Self { - nodes: Vec::new(), + nodes: ExecNodes { + items: Vec::new(), + pos: 0, + executing: 0, + }, result: ExecResult::default(), table_codec: TableCodec::default(), - projection_tmp: Vec::new(), context: None, transaction: std::ptr::null_mut(), runtime_probe_stack: Vec::new(), @@ -440,37 +507,6 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { } } -pub(crate) fn with_projection_tmp_value<'a, T: Transaction + 'a>( - 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> { - exec_arena.with_projection_tmp(|exec_arena, projection_tmp| { - { - 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(plan_arena, Some(tuple))?); - } - } - - match projection_tmp.len() { - 0 => {} - 1 => { - let value = projection_tmp.pop().expect("projection has one value"); - f(exec_arena, value)?; - } - _ => { - let value = DataValue::Tuple(std::mem::take(projection_tmp), false); - f(exec_arena, value)?; - } - } - Ok(()) - }) -} - impl<'a, T: Transaction + 'a> ExecArena<'a, T> { pub(crate) fn init_context(&mut self, context: ExecutionContext<'a>, transaction: &'a T) { if let Some(current) = &self.context { @@ -483,9 +519,7 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { } pub(crate) fn push(&mut self, node: ExecNode<'a, T>) -> ExecId { - let id = self.nodes.len(); - self.nodes.push(node); - id + self.nodes.push(node) } pub(crate) fn push_ddl_apply(&mut self, apply: DDLApply) { @@ -520,7 +554,7 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { pub(crate) fn local_state<'b>( &'b mut self, - plan_arena: &'b PlanArena<'a>, + plan_arena: &'b (dyn MetaArena + 'a), ) -> ExecArenaLocalState<'b, 'a, T> { let context = *self .context @@ -565,6 +599,7 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { debug_assert!(self.runtime_probe_stack.is_empty()); debug_assert!(self.ddl_apply.is_empty()); self.nodes.clear(); + self.result.tuple = Tuple::default(); self.result.status = None; self.recursive_input = None; } @@ -579,17 +614,40 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { &mut self.result.tuple } - #[inline] - pub(crate) fn with_projection_tmp( + pub(crate) fn materialize_tuple(&mut self) -> Tuple { + std::mem::take(&mut self.result.tuple) + } + + pub(crate) fn rewrite( &mut self, - f: impl FnOnce(&mut Self, &mut Vec) -> Result, - ) -> Result { - let mut projection_tmp = std::mem::take(&mut self.projection_tmp); - projection_tmp.clear(); - let ret = f(self, &mut projection_tmp); - projection_tmp.clear(); - self.projection_tmp = projection_tmp; - ret + exprs: &[E], + arena: &dyn MetaArena, + input: Option<&dyn TupleLike>, + ) -> Result<(), DatabaseError> { + let values = &mut self.result.tuple.values; + let base = values.len(); + values.reserve(exprs.len()); + + for expr in exprs { + let value = { + let input_values = &values[..base]; + let current: &dyn TupleLike = input.unwrap_or(&input_values); + expr.expression(arena) + .eval(arena, Some(current)) + .map(|value| value.into_owned()) + }; + match value { + Ok(value) => values.push(value), + Err(error) => { + values.truncate(base); + return Err(error); + } + } + } + + values.rotate_left(base); + values.truncate(exprs.len()); + Ok(()) } #[inline] @@ -611,12 +669,19 @@ impl<'a, T: Transaction + 'a> ExecArena<'a, T> { pub(crate) fn next_tuple( &mut self, id: ExecId, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), ) -> Result { self.result.status = None; - let mut node = std::mem::replace(&mut self.nodes[id], ExecNode::Empty); - let result = node.next_tuple(self, plan_arena); - self.nodes[id] = node; + let slot = &self.nodes.items[id] as *const std::cell::RefCell>; + self.nodes.executing += 1; + // SAFETY: push/clear cannot relocate or destroy slots while executing. + // Access to nodes always goes through RefCell: recursive calls and + // subtree rebuilds cannot borrow/overwrite an active node. The payload + // is behind UnsafeCell, separate from the arena's own mutable state. + // On unwind the borrow is released; executing remains nonzero, preventing + // relocation even if the caller catches the panic (the arena is poisoned). + let result = unsafe { (&*slot).borrow_mut().next_tuple(self, plan_arena) }; + self.nodes.executing -= 1; result?; match self.result.status.unwrap_or(ExecStatus::End) { @@ -632,7 +697,7 @@ pub(crate) trait ReadExecutor<'a, T: Transaction + 'a>: Sized { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId; @@ -644,7 +709,7 @@ pub(crate) trait WriteExecutor<'a, T: Transaction + 'a>: Sized { fn into_executor( input: Self::Input, arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), cache: ExecutionContext<'_>, transaction: &T, ) -> ExecId; @@ -652,7 +717,7 @@ pub(crate) trait WriteExecutor<'a, T: Transaction + 'a>: Sized { pub(crate) fn build_read<'a, T>( arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), plan: LogicalPlan, cache: ExecutionContext<'_>, transaction: &T, @@ -746,7 +811,7 @@ where } Operator::MarkApply(op) => { let (left, right) = childrens.pop_twins(); - as ReadExecutor<'a, T>>::into_executor( + >::into_executor( (op, left, right), arena, plan_arena, @@ -952,7 +1017,7 @@ where pub(crate) fn build_write<'a, T>( arena: &mut ExecArena<'a, T>, - plan_arena: &mut PlanArena<'a>, + plan_arena: &mut (dyn MetaArena + 'a), plan: LogicalPlan, cache: ExecutionContext<'a>, transaction: &'a mut T, @@ -1269,3 +1334,40 @@ mod test_utils { pub(crate) use test_utils::{ empty_context, execute, execute_input, execute_input_mut, execute_mut, try_collect, }; + +#[cfg(test)] +mod test { + use super::*; + use crate::storage::memory::MemoryTransaction; + use std::panic::{catch_unwind, AssertUnwindSafe}; + + #[test] + fn active_nodes_cannot_be_overwritten_or_relocated() { + let table_arena = crate::planner::TableArenaCell::default(); + let mut plan_arena = crate::planner::PlanArena::new(&table_arena); + let mut arena = ExecArena::<'_, MemoryTransaction>::new(); + arena.push(ExecNode::Dummy(Dummy::default())); + let slot = + &arena.nodes.items[0] as *const std::cell::RefCell>; + arena.nodes.executing = 1; + // Same access pattern as next_tuple; mutations must fail before changing storage. + let active = unsafe { (&*slot).borrow_mut() }; + assert!(catch_unwind(AssertUnwindSafe(|| { + arena.push(ExecNode::Dummy(Dummy::default())); + })) + .is_err()); + assert!(catch_unwind(AssertUnwindSafe(|| arena.nodes.clear())).is_err()); + arena.nodes.seek(0); + assert!(catch_unwind(AssertUnwindSafe(|| { + arena.push(ExecNode::Dummy(Dummy::default())); + })) + .is_err()); + assert!(catch_unwind(AssertUnwindSafe(|| { + arena.next_tuple(0, &mut plan_arena).unwrap(); + })) + .is_err()); + drop(active); + assert_eq!(arena.nodes.items.len(), 1); + assert_eq!(arena.nodes.items.as_ptr(), slot); + } +} diff --git a/src/execution/spill/codec.rs b/src/execution/spill/codec.rs index 4e000688..20e618a3 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::MetaArena; use crate::planner::PlanArena; use crate::types::tuple::Tuple; use crate::types::value::DataValue; @@ -30,11 +31,16 @@ impl SortRow { pub(crate) fn new( sort_fields: &[SortField], tuple: Tuple, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result { let sort_values = sort_fields .iter() - .map(|field| arena.expression(field.expr).eval(arena, Some(&tuple))) + .map(|field| { + arena + .expression(field.expr) + .eval(arena, Some(&tuple)) + .map(|v| v.into_owned()) + }) .collect::>()?; Ok(Self { sort_values, tuple }) } @@ -187,7 +193,7 @@ impl SpillCodec for Tuple { fn estimated_dynamic_value_size(value: &DataValue) -> usize { match value { DataValue::Utf8 { value, .. } => value.capacity(), - DataValue::Tuple(values, _) => values + DataValue::Tuple(values) => values .capacity() .saturating_mul(size_of::()) .saturating_add( diff --git a/src/execution/spill/mod.rs b/src/execution/spill/mod.rs index b79870e3..1fbb91eb 100644 --- a/src/execution/spill/mod.rs +++ b/src/execution/spill/mod.rs @@ -759,13 +759,10 @@ mod tests { let none = Option::::None; assert!(some.estimated_size() > none.estimated_size()); - let nested = DataValue::Tuple( - vec![ - DataValue::new_utf8("outer".to_string()), - DataValue::Tuple(vec![DataValue::new_utf8("inner".to_string())], false), - ], - false, - ); + let nested = DataValue::Tuple(vec![ + DataValue::new_utf8("outer".to_string()), + DataValue::Tuple(vec![DataValue::new_utf8("inner".to_string())]), + ]); assert!(nested.estimated_size() > std::mem::size_of::()); Ok(()) } diff --git a/src/expression/eq_col.rs b/src/expression/eq_col.rs index 0594ac8b..af814950 100644 --- a/src/expression/eq_col.rs +++ b/src/expression/eq_col.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -20,23 +22,28 @@ 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::planner::ExprRef; +use crate::planner::MetaArena; 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 { +pub(super) fn eq_ignore_colref_pos( + lhs: ExprRef, + rhs: ExprRef, + arena: &(dyn MetaArena + '_), +) -> bool { EqIgnoreColRefPosVisitor::equals(lhs, rhs, arena) } struct EqIgnoreColRefPosVisitor<'a, 'arena> { rhs: ExprRef, - arena: &'a PlanArena<'arena>, + arena: &'a (dyn MetaArena + 'arena), equal: bool, } impl<'a, 'arena> EqIgnoreColRefPosVisitor<'a, 'arena> { - fn equals(lhs: ExprRef, rhs: ExprRef, arena: &'a PlanArena<'arena>) -> bool { + fn equals(lhs: ExprRef, rhs: ExprRef, arena: &'a (dyn MetaArena + 'arena)) -> bool { let mut visitor = Self { rhs, arena, @@ -66,8 +73,8 @@ impl<'a, 'arena> EqIgnoreColRefPosVisitor<'a, 'arena> { } } -impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { - fn visit(&mut self, lhs: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { +impl ExprVisitor for EqIgnoreColRefPosVisitor<'_, '_> { + fn visit(&mut self, lhs: ExprRef, arena: &(dyn MetaArena + '_)) -> Result<(), DatabaseError> { let lhs = lhs.unpack_alias(arena); self.rhs = self.rhs.unpack_alias(arena); if lhs == self.rhs { @@ -95,7 +102,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_expr: ExprRef, lhs_ty: &LogicalType, lhs_evaluator: Option<&CastEvaluatorRef>, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::TypeCast { @@ -116,7 +123,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { &mut self, lhs_negated: bool, lhs_expr: ExprRef, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::IsNull { @@ -134,7 +141,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_expr: ExprRef, lhs_evaluator: Option<&UnaryEvaluatorRef>, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Unary { @@ -160,7 +167,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_right: ExprRef, lhs_evaluator: Option<&BinaryEvaluatorRef>, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Binary { @@ -187,7 +194,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_kind: &AggKind, lhs_args: &[ExprRef], lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::AggCall { @@ -209,7 +216,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { fn visit_window( &mut self, lhs: &WindowCall, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::WindowCall(rhs) => { @@ -236,7 +243,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_negated: bool, lhs_expr: ExprRef, lhs_args: &[ExprRef], - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::In { @@ -259,7 +266,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_expr: ExprRef, lhs_left: ExprRef, lhs_right: ExprRef, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Between { @@ -283,7 +290,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_expr: ExprRef, lhs_for: Option, lhs_from: Option, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::SubString { @@ -304,7 +311,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { &mut self, lhs_expr: ExprRef, lhs_in: ExprRef, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Position { @@ -324,7 +331,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_expr: ExprRef, lhs_what: Option, lhs_where: Option<&TrimWhereField>, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Trim { @@ -349,7 +356,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { fn visit_tuple( &mut self, lhs: &[ExprRef], - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = matches!(self.rhs(), ScalarExpression::Tuple(rhs) if self.refs_equal(lhs, rhs)); @@ -359,7 +366,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { fn visit_scala_function( &mut self, lhs: &ScalarFunction, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = matches!( self.rhs(), @@ -372,7 +379,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { fn visit_table_function( &mut self, lhs: &TableFunction, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = matches!( self.rhs(), @@ -388,7 +395,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_left: ExprRef, lhs_right: ExprRef, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::If { @@ -412,7 +419,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_left: ExprRef, lhs_right: ExprRef, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::IfNull { @@ -434,7 +441,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_left: ExprRef, lhs_right: ExprRef, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::NullIf { @@ -455,7 +462,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { &mut self, lhs_exprs: &[ExprRef], lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::Coalesce { @@ -473,7 +480,7 @@ impl ExprVisitor> for EqIgnoreColRefPosVisitor<'_, '_> { lhs_pairs: &[(ExprRef, ExprRef)], lhs_else: Option, lhs_ty: &LogicalType, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.equal = match self.rhs() { ScalarExpression::CaseWhen { @@ -513,7 +520,7 @@ mod tests { use crate::planner::TableArenaCell; fn assert_case( - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), lhs: ScalarExpression, rhs: ScalarExpression, different: ScalarExpression, diff --git a/src/expression/evaluator.rs b/src/expression/evaluator.rs index 391c4643..aff9f0e7 100644 --- a/src/expression/evaluator.rs +++ b/src/expression/evaluator.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -15,7 +17,8 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::ScalarFunction; use crate::expression::{AliasType, BinaryOperator, ScalarExpression, TrimWhereField}; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::evaluator::binary_create; use crate::types::tuple::TupleLike; use crate::types::value::{DataValue, Utf8Type}; @@ -26,36 +29,41 @@ use std::cmp::Ordering; macro_rules! eval_to_num { ($num_expr:expr, $arena:expr, $tuple:expr) => { - if let Some(num_i32) = $arena - .expression(*$num_expr) - .eval($arena, $tuple)? - .cast(&LogicalType::Integer)? - .i32() + if let Some(num_i32) = cast_cow( + $arena.expression(*$num_expr).eval($arena, $tuple)?, + &LogicalType::Integer, + )? + .i32() { num_i32 } else { - return Ok(DataValue::Null); + return Ok(Cow::Owned(DataValue::Null)); } }; } impl ScalarExpression { - pub fn eval( - &self, - arena: &PlanArena<'_>, - tuple: Option, - ) -> Result { + pub fn eval<'a>( + &'a self, + arena: &'a (dyn MetaArena + '_), + tuple: Option<&'a dyn TupleLike>, + ) -> Result, DatabaseError> { match self { - ScalarExpression::Constant(val) => Ok(val.clone()), + ScalarExpression::Constant(val) => match val { + DataValue::Parameter { id, .. } => { + Err(DatabaseError::parameter_not_found(format!("${id}"))) + } + val => Ok(Cow::Borrowed(val)), + }, ScalarExpression::ColumnRef { position, .. } => { let Some(tuple) = tuple else { - return Ok(DataValue::Null); + return Ok(Cow::Owned(DataValue::Null)); }; - Ok(tuple.value_at(*position).clone()) + Ok(Cow::Borrowed(tuple.value_at(*position))) } ScalarExpression::Alias { expr, alias } => { let Some(tuple) = tuple else { - return Ok(DataValue::Null); + return Ok(Cow::Owned(DataValue::Null)); }; if let AliasType::Expr(inner_expr) = alias { arena.expression(*inner_expr).eval(arena, Some(tuple)) @@ -68,7 +76,7 @@ impl ScalarExpression { } => { let value = arena.expression(*expr).eval(arena, tuple)?; if let Some(evaluator) = evaluator { - evaluator.eval(&value) + evaluator.eval(&value).map(Cow::Owned) } else { Ok(value) } @@ -86,13 +94,14 @@ impl ScalarExpression { .as_ref() .ok_or(DatabaseError::EvaluatorNotFound)? .binary_eval(&left, &right) + .map(Cow::Owned) } ScalarExpression::IsNull { expr, negated } => { let mut is_null = arena.expression(*expr).eval(arena, tuple)?.is_null(); if *negated { is_null = !is_null; } - Ok(DataValue::Boolean(is_null)) + Ok(Cow::Owned(DataValue::Boolean(is_null))) } ScalarExpression::In { expr, @@ -101,7 +110,7 @@ impl ScalarExpression { } => { let value = arena.expression(*expr).eval(arena, tuple)?; if value.is_null() { - return Ok(DataValue::Null); + return Ok(Cow::Owned(DataValue::Null)); } let mut matched = false; @@ -120,11 +129,11 @@ impl ScalarExpression { } if matched { - Ok(DataValue::Boolean(!negated)) + Ok(Cow::Owned(DataValue::Boolean(!negated))) } else if saw_null { - Ok(DataValue::Null) + Ok(Cow::Owned(DataValue::Null)) } else { - Ok(DataValue::Boolean(*negated)) + Ok(Cow::Owned(DataValue::Boolean(*negated))) } } ScalarExpression::Unary { @@ -132,10 +141,12 @@ impl ScalarExpression { } => { let value = arena.expression(*expr).eval(arena, tuple)?; - Ok(evaluator - .as_ref() - .ok_or(DatabaseError::EvaluatorNotFound)? - .unary_eval(&value)) + Ok(Cow::Owned( + evaluator + .as_ref() + .ok_or(DatabaseError::EvaluatorNotFound)? + .unary_eval(&value), + )) } ScalarExpression::AggCall { .. } => { unreachable!("must use `NormalizationRuleImpl::ExpressionRemapper`") @@ -155,26 +166,24 @@ impl ScalarExpression { value.partial_cmp(&right).map(Ordering::is_le), ) { (Some(true), Some(true)) => true, - (None, _) | (_, None) => return Ok(DataValue::Null), + (None, _) | (_, None) => return Ok(Cow::Owned(DataValue::Null)), _ => false, }; if *negated { is_between = !is_between; } - Ok(DataValue::Boolean(is_between)) + Ok(Cow::Owned(DataValue::Boolean(is_between))) } ScalarExpression::SubString { expr, for_expr, from_expr, } => { - if let Some(mut string) = arena - .expression(*expr) - .eval(arena, tuple)? - .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? - .utf8() - .map(String::from) - { + let value = cast_cow( + arena.expression(*expr).eval(arena, tuple)?, + &LogicalType::Varchar(None, CharLengthUnits::Characters), + )?; + if let Some(mut string) = value.utf8().map(String::from) { if let Some(from_expr) = from_expr { let mut from = eval_to_num!(from_expr, arena, tuple).saturating_sub(1); let len_i = string.len() as i32; @@ -183,7 +192,7 @@ impl ScalarExpression { from += len_i + 1; } if from > len_i { - return Ok(DataValue::Null); + return Ok(Cow::Owned(DataValue::Null)); } string = string.split_off(from as usize); } @@ -193,77 +202,71 @@ impl ScalarExpression { let _ = string.split_off(for_i); } - Ok(DataValue::Utf8 { + Ok(Cow::Owned(DataValue::Utf8 { value: string, ty: Utf8Type::Variable(None), unit: CharLengthUnits::Characters, - }) + })) } else { - Ok(DataValue::Null) + Ok(Cow::Owned(DataValue::Null)) } } ScalarExpression::Position { expr, in_expr } => { - let unpack = |expr: ExprRef| -> Result { - Ok(arena - .expression(expr) - .eval(arena, tuple)? - .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? + let varchar = LogicalType::Varchar(None, CharLengthUnits::Characters); + let pattern = cast_cow(arena.expression(*expr).eval(arena, tuple)?, &varchar)?; + let string = cast_cow(arena.expression(*in_expr).eval(arena, tuple)?, &varchar)?; + Ok(Cow::Owned(DataValue::Int32( + string .utf8() - .map(String::from) - .unwrap_or("".to_owned())) - }; - let pattern = unpack(*expr)?; - let str = unpack(*in_expr)?; - Ok(DataValue::Int32( - str.find(&pattern).map(|pos| pos as i32 + 1).unwrap_or(0), - )) + .unwrap_or("") + .find(pattern.utf8().unwrap_or("")) + .map(|pos| pos as i32 + 1) + .unwrap_or(0), + ))) } ScalarExpression::Trim { expr, trim_what_expr, trim_where, } => { - if let Some(string) = arena - .expression(*expr) - .eval(arena, tuple)? - .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? - .utf8() - { + let value = cast_cow( + arena.expression(*expr).eval(arena, tuple)?, + &LogicalType::Varchar(None, CharLengthUnits::Characters), + )?; + if let Some(string) = value.utf8() { let mut trim_what = String::from(" "); if let Some(trim_what_expr) = trim_what_expr { - trim_what = arena - .expression(*trim_what_expr) - .eval(arena, tuple)? - .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))? - .utf8() - .map(String::from) - .unwrap_or_default(); + let value = cast_cow( + arena.expression(*trim_what_expr).eval(arena, tuple)?, + &LogicalType::Varchar(None, CharLengthUnits::Characters), + )?; + trim_what = value.utf8().unwrap_or("").to_owned(); } let string_trimmed = trim_string(string, &trim_what, *trim_where); - Ok(DataValue::Utf8 { + Ok(Cow::Owned(DataValue::Utf8 { value: string_trimmed, ty: Utf8Type::Variable(None), unit: CharLengthUnits::Characters, - }) + })) } else { - Ok(DataValue::Null) + Ok(Cow::Owned(DataValue::Null)) } } ScalarExpression::Tuple(exprs) => { let mut values = Vec::with_capacity(exprs.len()); for expr in exprs { - values.push(arena.expression(*expr).eval(arena, tuple)?); + values.push(arena.expression(*expr).eval(arena, tuple)?.into_owned()); } - Ok(DataValue::Tuple(values, false)) + Ok(Cow::Owned(DataValue::Tuple(values))) } ScalarExpression::ScalaFunction(ScalarFunction { inner, args, .. }) => { let value = match tuple { - Some(tuple) => inner.eval(args, arena, Some(&tuple as &dyn TupleLike))?, + Some(tuple) => inner.eval(args, arena, Some(tuple))?, None => inner.eval(args, arena, None)?, }; - value.cast(inner.return_type()) + value.cast(inner.return_type()).map(Cow::Owned) } ScalarExpression::Empty => unreachable!(), ScalarExpression::If { @@ -273,9 +276,9 @@ impl ScalarExpression { ty, } => { if arena.expression(*condition).eval(arena, tuple)?.is_true()? { - arena.expression(*left_expr).eval(arena, tuple)?.cast(ty) + cast_cow(arena.expression(*left_expr).eval(arena, tuple)?, ty) } else { - arena.expression(*right_expr).eval(arena, tuple)?.cast(ty) + cast_cow(arena.expression(*right_expr).eval(arena, tuple)?, ty) } } ScalarExpression::IfNull { @@ -288,7 +291,7 @@ impl ScalarExpression { if value.is_null() { value = arena.expression(*right_expr).eval(arena, tuple)?; } - value.cast(ty) + cast_cow(value, ty) } ScalarExpression::NullIf { left_expr, @@ -298,9 +301,9 @@ impl ScalarExpression { let mut value = arena.expression(*left_expr).eval(arena, tuple)?; if arena.expression(*right_expr).eval(arena, tuple)? == value { - value = DataValue::Null; + value = Cow::Owned(DataValue::Null); } - value.cast(ty) + cast_cow(value, ty) } ScalarExpression::Coalesce { exprs, ty } => { let mut value = None; @@ -313,7 +316,7 @@ impl ScalarExpression { break; } } - value.unwrap_or(DataValue::Null).cast(ty) + cast_cow(value.unwrap_or(Cow::Owned(DataValue::Null)), ty) } ScalarExpression::CaseWhen { operand_expr, @@ -331,7 +334,7 @@ impl ScalarExpression { 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)?; + when_value = cast_cow(when_value, &ty)?; let evaluator = binary_create(Cow::Owned(ty), BinaryOperator::Eq)?; evaluator .binary_eval(operand_value, &when_value)? @@ -349,7 +352,7 @@ impl ScalarExpression { result = Some(arena.expression(*expr).eval(arena, tuple)?); } } - result.unwrap_or(DataValue::Null).cast(ty) + cast_cow(result.unwrap_or(Cow::Owned(DataValue::Null)), ty) } ScalarExpression::TableFunction(_) => unreachable!(), ScalarExpression::WindowCall(_) => Err(DatabaseError::UnsupportedStmt( @@ -359,6 +362,17 @@ impl ScalarExpression { } } +fn cast_cow<'a>( + value: Cow<'a, DataValue>, + ty: &LogicalType, +) -> Result, DatabaseError> { + if value.logical_type() == *ty { + Ok(value) + } else { + value.into_owned().cast(ty).map(Cow::Owned) + } +} + fn trim_string(value: &str, trim_what: &str, trim_where: Option) -> String { if trim_what.is_empty() { return value.to_string(); @@ -387,18 +401,16 @@ fn trim_string(value: &str, trim_what: &str, trim_where: Option) #[cfg(test)] mod tests { use super::*; + use crate::planner::test::PlanArenaTestExt; fn const_in( - arena: &mut PlanArena, + 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(); + let expr = arena.alloc_expression(expr.into()); + let args = arena.alloc_expressions(args); arena.alloc_expression(ScalarExpression::In { negated, expr, @@ -406,6 +418,114 @@ mod tests { }) } + #[test] + fn eval_borrows_leaf_values_and_owns_binary_results() -> Result<(), DatabaseError> { + use crate::types::evaluator::binary_create; + use crate::types::tuple::Tuple; + + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let value = DataValue::Int64(42); + let constant = arena.alloc_expression(ScalarExpression::Constant(value.clone())); + let column_ref = arena.alloc_column(crate::catalog::ColumnCatalog::new( + "value".to_owned(), + true, + crate::catalog::ColumnDesc::new(LogicalType::Bigint, None, false, None)?, + )); + let column = arena.alloc_expression(ScalarExpression::ColumnRef { + column: column_ref, + position: 0, + }); + let sum = arena.alloc_expression(ScalarExpression::Binary { + op: BinaryOperator::Plus, + left_expr: constant, + right_expr: column, + evaluator: Some(binary_create( + Cow::Owned(LogicalType::Bigint), + BinaryOperator::Plus, + )?), + ty: LogicalType::Bigint, + }); + let tuple = Tuple::new(None, vec![value.clone()]); + assert!(matches!( + arena.expression(constant).eval(&arena, Some(&tuple)), + Ok(Cow::Borrowed(_)) + )); + assert!(matches!( + arena.expression(column).eval(&arena, Some(&tuple)), + Ok(Cow::Borrowed(_)) + )); + assert_eq!( + arena.expression(sum).eval(&arena, Some(&tuple))?, + Cow::Owned(DataValue::Int64(84)) + ); + + Ok(()) + } + + #[test] + fn eval_borrows_passthrough_branches_and_owns_casts() -> Result<(), DatabaseError> { + use crate::types::evaluator::cast_create; + use crate::types::tuple::Tuple; + + let table_arena = crate::planner::TableArenaCell::default(); + let mut arena = PlanArena::new(&table_arena); + let value = DataValue::Int32(7); + let column_ref = arena.alloc_column(crate::catalog::ColumnCatalog::new( + "value".to_owned(), + true, + crate::catalog::ColumnDesc::new(LogicalType::Integer, None, false, None)?, + )); + let column = arena.alloc_expression(ScalarExpression::ColumnRef { + column: column_ref, + position: 0, + }); + let alias = arena.alloc_expression(ScalarExpression::Alias { + expr: column, + alias: AliasType::Name("alias".to_owned()), + }); + let no_op_cast = arena.alloc_expression(ScalarExpression::TypeCast { + expr: column, + ty: LogicalType::Integer, + evaluator: None, + }); + let condition = + arena.alloc_expression(ScalarExpression::Constant(DataValue::Boolean(true))); + let null = arena.alloc_expression(ScalarExpression::Constant(DataValue::Null)); + let if_expr = arena.alloc_expression(ScalarExpression::If { + condition, + left_expr: column, + right_expr: null, + ty: LogicalType::Integer, + }); + let if_null = arena.alloc_expression(ScalarExpression::IfNull { + left_expr: null, + right_expr: column, + ty: LogicalType::Integer, + }); + let coalesce = arena.alloc_expression(ScalarExpression::Coalesce { + exprs: vec![null, column], + ty: LogicalType::Integer, + }); + let cast = arena.alloc_expression(ScalarExpression::TypeCast { + expr: column, + ty: LogicalType::Bigint, + evaluator: Some(cast_create(&LogicalType::Integer, &LogicalType::Bigint)?), + }); + let tuple = Tuple::new(None, vec![value.clone()]); + for expr in [alias, no_op_cast, if_expr, if_null, coalesce] { + match arena.expression(expr).eval(&arena, Some(&tuple))? { + Cow::Borrowed(actual) => assert!(std::ptr::eq(actual, &tuple.values[0])), + Cow::Owned(_) => panic!("pass-through expression must borrow the column"), + } + } + assert_eq!( + arena.expression(cast).eval(&arena, Some(&tuple))?, + Cow::Owned(DataValue::Int64(7)) + ); + Ok(()) + } + #[test] fn in_eval_matches_even_if_null_appears_first() -> Result<(), DatabaseError> { let table_arena = crate::planner::TableArenaCell::default(); @@ -418,7 +538,7 @@ mod tests { ); assert_eq!( - arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + arena.expression(expr).eval(&arena, None)?.into_owned(), DataValue::Boolean(true) ); Ok(()) @@ -436,7 +556,7 @@ mod tests { ); assert_eq!( - arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + arena.expression(expr).eval(&arena, None)?.into_owned(), DataValue::Null ); Ok(()) @@ -454,7 +574,7 @@ mod tests { ); assert_eq!( - arena.expression(expr).eval::<&[DataValue]>(&arena, None)?, + arena.expression(expr).eval(&arena, None)?.into_owned(), DataValue::Boolean(false) ); Ok(()) diff --git a/src/expression/function/scala.rs b/src/expression/function/scala.rs index 51d493a2..c2e755da 100644 --- a/src/expression/function/scala.rs +++ b/src/expression/function/scala.rs @@ -14,7 +14,8 @@ use crate::errors::DatabaseError; use crate::expression::function::FunctionSummary; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -64,7 +65,7 @@ pub trait ScalarFunctionImpl: Debug + Send + Sync { fn eval( &self, args: &[ExprRef], - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), tuple: Option<&dyn TupleLike>, ) -> Result; diff --git a/src/expression/function/table.rs b/src/expression/function/table.rs index e01491ad..ff746410 100644 --- a/src/expression/function/table.rs +++ b/src/expression/function/table.rs @@ -14,7 +14,8 @@ use crate::errors::DatabaseError; use crate::expression::function::FunctionSummary; -use crate::planner::{ExprRef, PlanArena, TableArena}; +use crate::planner::MetaArena; +use crate::planner::{ExprRef, TableArena}; use crate::types::tuple::{Schema, Tuple}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt::Debug; @@ -62,7 +63,7 @@ pub trait TableFunctionImpl: Debug + Send + Sync { fn eval( &self, args: &[ExprRef], - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result>>, DatabaseError>; fn summary(&self) -> &FunctionSummary; diff --git a/src/expression/mod.rs b/src/expression/mod.rs index 3e79bb45..58663675 100644 --- a/src/expression/mod.rs +++ b/src/expression/mod.rs @@ -20,7 +20,7 @@ use crate::expression::function::table::TableFunction; use crate::expression::visitor::{walk_expr, ExprVisitor}; use crate::expression::visitor_mut::ExprVisitorMut; use crate::planner::operator::sort::SortField; -use crate::planner::{Explain, ExprRef, MetaArena, PlanArena}; +use crate::planner::{Explain, ExprRef, MetaArena}; use crate::types::evaluator::{ binary_create, cast_create, unary_create, BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef, @@ -283,7 +283,7 @@ impl ExprVisitorMut for BindEvaluator { expr: &mut ExprRef, ty: &mut LogicalType, evaluator: &mut Option, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; let from = expr.return_type(arena); @@ -302,7 +302,7 @@ impl ExprVisitorMut for BindEvaluator { expr: &mut ExprRef, evaluator: &mut Option, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; @@ -329,7 +329,7 @@ impl ExprVisitorMut for BindEvaluator { right_expr: &mut ExprRef, evaluator: &mut Option, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(left_expr, arena)?; self.visit(right_expr, arena)?; @@ -351,14 +351,14 @@ pub struct HasCountStar { pub value: bool, } -impl ExprVisitor> for HasCountStar { +impl ExprVisitor for HasCountStar { fn visit_agg( &mut self, _distinct: bool, _kind: &AggKind, args: &[ExprRef], _ty: &LogicalType, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if args.len() == 1 { if let ScalarExpression::Constant(value) = arena.expression(args[0]) { @@ -368,7 +368,7 @@ impl ExprVisitor> for HasCountStar { Ok(()) } - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + fn visit(&mut self, expr: ExprRef, arena: &(dyn MetaArena + '_)) -> Result<(), DatabaseError> { if !self.value { walk_expr(self, expr, arena)?; } @@ -377,19 +377,19 @@ impl ExprVisitor> for HasCountStar { } pub trait TypeCast: Sized { - fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType>; + fn return_type<'a>(&'a self, arena: &'a (dyn MetaArena + '_)) -> Cow<'a, LogicalType>; fn into_expr( self, ty: LogicalType, evaluator: CastEvaluatorRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Self; fn type_cast( self, ty: Cow<'_, LogicalType>, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result { let from = self.return_type(arena); if from.as_ref() == ty.as_ref() { @@ -401,7 +401,7 @@ pub trait TypeCast: Sized { } impl TypeCast for ScalarExpression { - fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType> { + fn return_type<'a>(&'a self, arena: &'a (dyn MetaArena + '_)) -> Cow<'a, LogicalType> { match self { ScalarExpression::Constant(value) => Cow::Owned(value.logical_type()), ScalarExpression::ColumnRef { column, .. } => { @@ -445,7 +445,7 @@ impl TypeCast for ScalarExpression { self, ty: LogicalType, evaluator: CastEvaluatorRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Self { ScalarExpression::TypeCast { expr: arena.alloc_expression(self), @@ -456,7 +456,7 @@ impl TypeCast for ScalarExpression { } impl TypeCast for ExprRef { - fn return_type<'a>(&'a self, arena: &'a PlanArena<'_>) -> Cow<'a, LogicalType> { + fn return_type<'a>(&'a self, arena: &'a (dyn MetaArena + '_)) -> Cow<'a, LogicalType> { arena.expression(*self).return_type(arena) } @@ -464,7 +464,7 @@ impl TypeCast for ExprRef { self, ty: LogicalType, evaluator: CastEvaluatorRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Self { arena.alloc_expression(ScalarExpression::TypeCast { expr: self, @@ -481,10 +481,10 @@ impl ScalarExpression { } impl Explain for ExprRef { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, arena: &dyn MetaArena, f: &mut fmt::Formatter<'_>) -> fmt::Result { fn write_exprs( exprs: &[ExprRef], - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>, ) -> fmt::Result { for (index, expr) in exprs.iter().enumerate() { @@ -731,20 +731,20 @@ impl ExprRef { SortField::from(self).nulls_last() } - pub(crate) fn eq_ignore_colref_pos(self, other: ExprRef, arena: &PlanArena) -> bool { + pub(crate) fn eq_ignore_colref_pos(self, other: ExprRef, arena: &(dyn MetaArena + '_)) -> bool { eq_col::eq_ignore_colref_pos(self, other, arena) } pub(crate) fn clone_expression( self, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> 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 { + pub fn unpack_alias(self, arena: &(impl MetaArena + ?Sized)) -> ExprRef { if let ScalarExpression::Alias { alias: AliasType::Expr(expr), .. @@ -758,25 +758,29 @@ impl ExprRef { } } - pub fn unpack_alias_ref<'a, A: MetaArena>(self, arena: &'a A) -> &'a ScalarExpression { + pub fn unpack_alias_ref<'a, A: MetaArena + ?Sized>(self, arena: &'a A) -> &'a ScalarExpression { arena.expression(self.unpack_alias(arena)) } pub fn any_referenced_column( self, - arena: &PlanArena, - mut predicate: impl FnMut(&PlanArena, &ColumnRef) -> bool, + arena: &(dyn MetaArena + '_), + mut predicate: impl FnMut(&(dyn MetaArena + '_), &ColumnRef) -> bool, ) -> Result { struct ColumnRefVisitor<'a, 'arena, F> { f: &'a mut F, any: bool, - arena: &'a PlanArena<'arena>, + arena: &'a (dyn MetaArena + 'arena), } - impl bool> ExprVisitor> + impl bool> ExprVisitor for ColumnRefVisitor<'_, '_, F> { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + fn visit( + &mut self, + expr: ExprRef, + arena: &(dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { if !self.any { walk_expr(self, expr, arena)?; } @@ -800,19 +804,23 @@ impl ExprRef { pub fn all_referenced_columns( self, - arena: &PlanArena, - mut predicate: impl FnMut(&PlanArena, &ColumnRef) -> bool, + arena: &(dyn MetaArena + '_), + mut predicate: impl FnMut(&(dyn MetaArena + '_), &ColumnRef) -> bool, ) -> Result { struct ColumnRefVisitor<'a, 'arena, F> { f: &'a mut F, all: bool, - arena: &'a PlanArena<'arena>, + arena: &'a (dyn MetaArena + 'arena), } - impl bool> ExprVisitor> + impl bool> ExprVisitor for ColumnRefVisitor<'_, '_, F> { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + fn visit( + &mut self, + expr: ExprRef, + arena: &(dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { if self.all { walk_expr(self, expr, arena)?; } @@ -834,12 +842,16 @@ impl ExprRef { Ok(visitor.all) } - pub fn has_agg_call(self, arena: &PlanArena<'_>) -> Result { + pub fn has_agg_call(self, arena: &(dyn MetaArena + '_)) -> Result { struct AggCallChecker { has_agg: bool, } - impl ExprVisitor> for AggCallChecker { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + impl ExprVisitor for AggCallChecker { + fn visit( + &mut self, + expr: ExprRef, + arena: &(dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { if self.has_agg { return Ok(()); } @@ -851,7 +863,7 @@ impl ExprRef { _kind: &AggKind, args: &[ExprRef], _ty: &LogicalType, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for arg in args { self.visit(*arg, arena)?; @@ -865,11 +877,15 @@ impl ExprRef { Ok(checker.has_agg) } - pub fn has_window_call(self, arena: &PlanArena<'_>) -> Result { + pub fn has_window_call(self, arena: &(dyn MetaArena + '_)) -> Result { struct WindowCallChecker(bool); - impl ExprVisitor> for WindowCallChecker { - fn visit(&mut self, expr: ExprRef, arena: &PlanArena<'_>) -> Result<(), DatabaseError> { + impl ExprVisitor for WindowCallChecker { + fn visit( + &mut self, + expr: ExprRef, + arena: &(dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { if !self.0 { walk_expr(self, expr, arena)?; } @@ -879,7 +895,7 @@ impl ExprRef { fn visit_window( &mut self, _window: &window::WindowCall, - _arena: &PlanArena<'_>, + _arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.0 = true; Ok(()) @@ -891,11 +907,11 @@ impl ExprRef { Ok(checker.0) } - pub fn output_name(self, arena: &PlanArena) -> String { + pub fn output_name(self, arena: &(dyn MetaArena + '_)) -> String { self.explain(arena).to_string() } - pub fn output_column_ref(self, arena: &mut PlanArena) -> ColumnRef { + pub fn output_column_ref(self, arena: &mut (dyn MetaArena + '_)) -> ColumnRef { match arena.expression(self) { ScalarExpression::ColumnRef { column, .. } => *column, ScalarExpression::Alias { @@ -1013,7 +1029,7 @@ mod test { use crate::expression::{AliasType, BinaryOperator, ScalarExpression, UnaryOperator}; use crate::function::current_date::CurrentDate; use crate::function::numbers::Numbers; - use crate::planner::{ExprRef, PlanArena, TableArenaCell}; + use crate::planner::{ExprRef, MetaArena, PlanArena, TableArenaCell}; use crate::serdes::{ReferenceDecodeContext, ReferenceSerialization, ReferenceTables}; use crate::storage::rocksdb::RocksStorage; use crate::storage::rocksdb::RocksTransaction; @@ -1033,7 +1049,7 @@ mod test { expr: ScalarExpression, drive: Option<&ReferenceDecodeContext<'_, RocksTransaction>>, reference_tables: &mut ReferenceTables, - arena: &mut PlanArena, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { let expr = arena.alloc_expression(expr); expr.encode(cursor, false, reference_tables, arena)?; diff --git a/src/expression/range_detacher.rs b/src/expression/range_detacher.rs index 0a65d2ef..964fd8d5 100644 --- a/src/expression/range_detacher.rs +++ b/src/expression/range_detacher.rs @@ -16,11 +16,13 @@ use crate::catalog::ColumnRef; use crate::errors::DatabaseError; use crate::expression::{BinaryOperator, ScalarExpression}; use crate::iter_ext::Itertools; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::index::IndexMetaRef; use crate::types::value::DataValue; use crate::types::{ColumnId, LogicalType}; use kite_sql_serde_macros::ReferenceSerialization; +use std::borrow::Borrow; use std::cmp::Ordering; use std::collections::Bound; use std::fmt::Formatter; @@ -54,10 +56,10 @@ impl DetachedPredicate { } } - fn combine_residuals( + fn combine_residuals( left: Option, right: Option, - arena: &mut PlanArena<'_>, + arena: &mut A, ) -> Option { match (left, right) { (Some(left), Some(right)) => Some(arena.alloc_expression(ScalarExpression::Binary { @@ -108,45 +110,78 @@ impl TreeNode { } } -fn build_tree(ranges: &[Range], current_level: usize) -> Option> { - fn build_subtree<'a>( - ranges: &'a [Range], - range: &'a Range, - current_level: usize, - ) -> Option> { - let value = match range { - Range::Eq(value) => value, - _ => return None, +fn build_tree(mut ranges: I) -> Option> +where + I: Iterator + Clone, + I::Item: Borrow, +{ + fn build_subtree(ranges: I, range: &Range) -> Option> + where + I: Iterator + Clone, + I::Item: Borrow, + { + let Range::Eq(value) = range else { + return None; }; - let mut child = TreeNode::new(Some(value)); - let subtree = build_tree(ranges, current_level + 1)?; - - if !subtree.children.is_empty() || current_level == ranges.len() - 1 { + let mut child = TreeNode::new(Some(value.clone())); + let is_last = ranges.clone().next().is_none(); + let subtree = build_tree(ranges)?; + if !subtree.children.is_empty() || is_last { child.add_child(subtree); } Some(child) } let mut root = TreeNode::new(None); - - if current_level < ranges.len() { - match &ranges[current_level] { + if let Some(range) = ranges.next() { + match range.borrow() { Range::SortedRanges(child_ranges) => { - for range in child_ranges.iter() { - root.children - .push(build_subtree(ranges, range, current_level)?); + for range in child_ranges { + root.children.push(build_subtree(ranges.clone(), range)?); } } - range => { - root.children - .push(build_subtree(ranges, range, current_level)?); - } + range => root.children.push(build_subtree(ranges, range)?), } } Some(root) } impl Range { + pub(crate) fn bind_parameters( + &mut self, + params: &[(usize, DataValue)], + ) -> Result<(), DatabaseError> { + match self { + Self::Scope { min, max } => { + for bound in [min, max] { + if let Bound::Included(value) | Bound::Excluded(value) = bound { + value.bind_parameters(params)?; + } + } + } + Self::Eq(value) => value.bind_parameters(params)?, + Self::SortedRanges(ranges) => { + for range in ranges { + range.bind_parameters(params)?; + } + } + Self::Dummy => {} + } + Ok(()) + } + + pub(crate) fn has_parameter(&self) -> bool { + match self { + Self::Scope { min, max } => [min, max].into_iter().any(|bound| match bound { + Bound::Included(value) | Bound::Excluded(value) => value.has_parameter(), + Bound::Unbounded => false, + }), + Self::Eq(value) => value.has_parameter(), + Self::SortedRanges(ranges) => ranges.iter().any(Self::has_parameter), + Self::Dummy => false, + } + } + pub(crate) fn only_eq(&self) -> bool { match self { Range::Eq(_) => true, @@ -155,49 +190,45 @@ impl Range { } } - pub(crate) fn combining_eqs(&self, eqs: &[Range]) -> Option { - #[allow(clippy::map_clone)] - fn merge_value(tuple: &[&DataValue], is_upper: bool, value: DataValue) -> DataValue { + pub(crate) fn combining_eqs(&self, eqs: I) -> Option + where + I: IntoIterator, + I::IntoIter: Clone, + I::Item: Borrow, + { + fn merge_value(tuple: &[DataValue], value: DataValue) -> DataValue { let mut merge_tuple = Vec::with_capacity(tuple.len() + 1); for value in tuple { merge_tuple.push((*value).clone()); } merge_tuple.push(value); - DataValue::Tuple(merge_tuple, is_upper) + DataValue::Tuple(merge_tuple) } - fn collect_tuple_range(result_ranges: &mut Vec, tuple: &[&DataValue], range: Range) { + fn collect_tuple_range(result_ranges: &mut Vec, tuple: &[DataValue], range: Range) { fn merge_value_on_bound( - tuple: &[&DataValue], - is_upper: bool, + tuple: &[DataValue], bound: Bound, ) -> Bound { match bound { - Bound::Included(v) => Bound::Included(merge_value(tuple, is_upper, v)), - Bound::Excluded(v) => Bound::Excluded(merge_value(tuple, !is_upper, v)), + Bound::Included(v) => Bound::Included(merge_value(tuple, v)), + Bound::Excluded(v) => Bound::Excluded(merge_value(tuple, v)), Bound::Unbounded => { if tuple.is_empty() { return Bound::Unbounded; } let values = tuple.iter().map(|v| (*v).clone()).collect_vec(); - // Excluding a lower equality prefix skips its entire key - // range when storage encodes the exclusive bound. Start at - // the prefix itself; the upper sentinel stays exclusive. - if is_upper { - Bound::Excluded(DataValue::Tuple(values, true)) - } else { - Bound::Included(DataValue::Tuple(values, false)) - } + Bound::Included(DataValue::Tuple(values)) } } } match range { Range::Scope { min, max } => result_ranges.push(Range::Scope { - min: merge_value_on_bound(tuple, false, min), - max: merge_value_on_bound(tuple, true, max), + min: merge_value_on_bound(tuple, min), + max: merge_value_on_bound(tuple, max), }), - Range::Eq(v) => result_ranges.push(Range::Eq(merge_value(tuple, false, v))), + Range::Eq(v) => result_ranges.push(Range::Eq(merge_value(tuple, v))), Range::Dummy => result_ranges.push(Range::Dummy), Range::SortedRanges(mut ranges) => { for range in &mut ranges { @@ -207,7 +238,7 @@ impl Range { } } - let node = build_tree(eqs, 0)?; + let node = build_tree(eqs.into_iter())?; let mut combinations = Vec::new(); node.enumeration(&mut Vec::new(), &mut combinations); @@ -222,7 +253,12 @@ impl Range { } pub trait RangeColumnMatcher { - fn matches(&self, table_name: &str, column_id: ColumnId, arena: &PlanArena<'_>) -> bool; + fn matches( + &self, + table_name: &str, + column_id: ColumnId, + arena: &A, + ) -> bool; } pub struct IndexRangeColumn { @@ -231,24 +267,94 @@ pub struct IndexRangeColumn { } impl RangeColumnMatcher for IndexRangeColumn { - fn matches(&self, table_name: &str, column_id: ColumnId, arena: &PlanArena<'_>) -> bool { + fn matches( + &self, + table_name: &str, + column_id: ColumnId, + arena: &A, + ) -> 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> { +pub struct RangeDetacher< + 'a, + M: RangeColumnMatcher = IndexRangeColumn, + A: MetaArena + ?Sized = dyn MetaArena + 'a, +> { column: M, - arena: &'a mut PlanArena<'p>, + arena: &'a mut A, } -impl<'a, 'p> RangeDetacher<'a, 'p, IndexRangeColumn> { - pub(crate) fn for_index( +impl<'a, A: MetaArena + ?Sized> RangeDetacher<'a, IndexRangeColumn, A> { + pub(crate) fn specialize_range( meta: IndexMetaRef, - position: usize, - arena: &'a mut PlanArena<'p>, - ) -> Self { + original: &Range, + predicate: ExprRef, + prefix_len: usize, + arena: &'a mut A, + ) -> Result, DatabaseError> { + let index = arena.index(meta); + if prefix_len >= index.column_ids.len() { + return Ok(None); + } + let composite = matches!(index.value_ty, LogicalType::Tuple(_)); + let Range::Scope { min, max } = original else { + return Ok(None); + }; + let prefix: &[DataValue] = if composite && prefix_len > 0 { + let ( + Bound::Included(DataValue::Tuple(lower)) | Bound::Excluded(DataValue::Tuple(lower)), + Bound::Included(DataValue::Tuple(upper)) | Bound::Excluded(DataValue::Tuple(upper)), + ) = (min, max) + else { + return Ok(None); + }; + let (Some(lower_prefix), Some(upper_prefix)) = + (lower.get(..prefix_len), upper.get(..prefix_len)) + else { + return Ok(None); + }; + if lower_prefix != upper_prefix + || (lower.len() == prefix_len && matches!(min, Bound::Excluded(_))) + || (upper.len() == prefix_len && matches!(max, Bound::Excluded(_))) + { + return Ok(None); + } + lower_prefix + } else { + &[] + }; + let Ok(Some(detached)) = Self::for_index(meta, prefix_len, arena).detach(predicate) else { + return Ok(None); + }; + // Only tighten a continuous range; leave equality, disjunction and empty + // ranges to the original lookup and residual Filter. + if !matches!(detached.range, Range::Scope { .. }) { + return Ok(None); + } + let constraint = if composite { + let Some(range) = detached + .range + .combining_eqs(prefix.iter().cloned().map(Range::Eq)) + else { + return Ok(None); + }; + range + } else { + detached.range + }; + Ok( + match Self::merge_binary(BinaryOperator::And, original.clone(), constraint) { + Ok(range) if !matches!(range, Range::Dummy) && &range != original => Some(range), + _ => None, + }, + ) + } + + pub(crate) fn for_index(meta: IndexMetaRef, position: usize, arena: &'a mut A) -> Self { Self { column: IndexRangeColumn { meta, position }, arena, @@ -256,7 +362,7 @@ impl<'a, 'p> RangeDetacher<'a, 'p, IndexRangeColumn> { } } -impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { +impl<'a, M: RangeColumnMatcher, A: MetaArena + ?Sized> RangeDetacher<'a, M, A> { pub(crate) fn detach( &mut self, expr: ExprRef, @@ -294,18 +400,27 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { let left = self.detach(left_expr)?; let right = self.detach(right_expr)?; 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) - else { - return Ok(None); - }; - let residual = DetachedPredicate::combine_residuals( - left_range.residual, - right_range.residual, - self.arena, - ); - (range, residual) + (Some(left), Some(right)) => { + let DetachedPredicate { + range: left_range, + residual: left_residual, + } = left; + let DetachedPredicate { + range: right_range, + residual: right_residual, + } = right; + + match Self::merge_binary(op, left_range, right_range) { + Ok(range) => { + let residual = DetachedPredicate::combine_residuals( + left_residual, + right_residual, + self.arena, + ); + (range, residual) + } + Err((left_range, _right_range)) => (left_range, Some(expr)), + } } (Some(detached), None) => { let residual = DetachedPredicate::combine_residuals( @@ -332,8 +447,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { 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 Ok(range) = Self::merge_binary(op, left.range, right.range) { return Ok(Some(DetachedPredicate::consumed(range))); } } @@ -406,7 +520,11 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { }) } - fn merge_binary(op: BinaryOperator, left_binary: Range, right_binary: Range) -> Option { + fn merge_binary( + op: BinaryOperator, + left_binary: Range, + right_binary: Range, + ) -> Result { fn process_exclude_bound_with_eq( bound: Bound, eq: &DataValue, @@ -423,11 +541,24 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { bound => bound, } } + + fn bounds_have_parameter(bounds: &[&Bound]) -> bool { + bounds.iter().any(|bound| match bound { + Bound::Included(value) | Bound::Excluded(value) => value.has_parameter(), + Bound::Unbounded => false, + }) + } + match (left_binary, right_binary) { - (Range::Dummy, binary) | (binary, Range::Dummy) => match op { - BinaryOperator::And => Some(Range::Dummy), - BinaryOperator::Or => Some(binary), - _ => None, + (Range::Dummy, binary) => match op { + BinaryOperator::And => Ok(Range::Dummy), + BinaryOperator::Or => Ok(binary), + _ => Err((Range::Dummy, binary)), + }, + (binary, Range::Dummy) => match op { + BinaryOperator::And => Ok(Range::Dummy), + BinaryOperator::Or => Ok(binary), + _ => Err((binary, Range::Dummy)), }, // e.g. c1 > 1 ? c1 < 2 ( @@ -440,17 +571,44 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { max: right_max, }, ) => match op { - BinaryOperator::And => Some(Self::and_scope_merge( - left_min, left_max, right_min, right_max, - )), - BinaryOperator::Or => Some(Self::or_scope_merge( - left_min, left_max, right_min, right_max, + BinaryOperator::And => { + Self::and_scope_merge(left_min, left_max, right_min, right_max) + } + BinaryOperator::Or => { + if bounds_have_parameter(&[&left_min, &left_max, &right_min, &right_max]) { + Err(( + Range::Scope { + min: left_min, + max: left_max, + }, + Range::Scope { + min: right_min, + max: right_max, + }, + )) + } else { + Ok(Self::or_scope_merge( + left_min, left_max, right_min, right_max, + )) + } + } + _ => Err(( + Range::Scope { + min: left_min, + max: left_max, + }, + Range::Scope { + min: right_min, + max: right_max, + }, )), - _ => None, }, // e.g. c1 > 1 ? c1 = 1 - (Range::Scope { min, max }, Range::Eq(eq)) - | (Range::Eq(eq), Range::Scope { min, max }) => { + (Range::Scope { min, max }, Range::Eq(eq)) => { + if bounds_have_parameter(&[&min, &max]) || eq.has_parameter() { + return Err((Range::Scope { min, max }, Range::Eq(eq))); + } + let unpack_bound = |bound_eq: Bound| match bound_eq { Bound::Included(val) | Bound::Excluded(val) => val, _ => unreachable!(), @@ -459,7 +617,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { BinaryOperator::And => { let bound_eq = Bound::Included(eq); let is_less = matches!( - Self::bound_compared(&bound_eq, &min, true).unwrap_or({ + Self::bound_compared(&bound_eq, &min, false, false).unwrap_or({ if matches!(min, Bound::Unbounded) { Ordering::Greater } else { @@ -471,24 +629,24 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { if is_less || matches!( - Self::bound_compared(&bound_eq, &max, false), + Self::bound_compared(&bound_eq, &max, true, true), Some(Ordering::Greater) ) { - return Some(Range::Dummy); + return Ok(Range::Dummy); } - Some(Range::Eq(unpack_bound(bound_eq))) + Ok(Range::Eq(unpack_bound(bound_eq))) } BinaryOperator::Or => { if eq.is_null() { - return Some(if matches!(min, Bound::Excluded(_)) { + return Ok(if matches!(min, Bound::Excluded(_)) { Range::SortedRanges(vec![Range::Eq(eq), Range::Scope { min, max }]) } else { Range::Scope { min, max } }); } let bound_eq = Bound::Excluded(eq); - let range = match Self::bound_compared(&bound_eq, &min, true) { + let range = match Self::bound_compared(&bound_eq, &min, false, false) { Some(Ordering::Less) => Range::SortedRanges(vec![ Range::Eq(unpack_bound(bound_eq)), Range::Scope { min, max }, @@ -501,7 +659,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { ), max, }, - _ => match Self::bound_compared(&bound_eq, &max, false) { + _ => match Self::bound_compared(&bound_eq, &max, true, true) { Some(Ordering::Greater) => Range::SortedRanges(vec![ Range::Scope { min, max }, Range::Eq(unpack_bound(bound_eq)), @@ -517,30 +675,41 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { _ => Range::Scope { min, max }, }, }; - Some(range) + Ok(range) } - _ => None, + _ => Err((Range::Scope { min, max }, Range::Eq(eq))), } } + // e.g. c1 = 1 ? c1 > 1 + (Range::Eq(eq), Range::Scope { min, max }) => { + Self::merge_binary(op, Range::Scope { min, max }, Range::Eq(eq)) + .map_err(|(right, left)| (left, right)) + } // e.g. c1 > 1 ? (c1 = 1 or c1 = 2) - (Range::Scope { min, max }, Range::SortedRanges(ranges)) - | (Range::SortedRanges(ranges), Range::Scope { min, max }) => { + (Range::Scope { min, max }, Range::SortedRanges(ranges)) => { + if bounds_have_parameter(&[&min, &max]) || ranges.iter().any(Range::has_parameter) { + return Err((Range::Scope { min, max }, Range::SortedRanges(ranges))); + } let merged_ranges = Self::extract_merge_ranges(op, Some(Range::Scope { min, max }), ranges, &mut 0); - - Some(Self::ranges2range(merged_ranges)) + Ok(Self::ranges2range(merged_ranges)) + } + // e.g. (c1 = 1 or c1 = 2) ? c1 > 1 + (Range::SortedRanges(ranges), Range::Scope { min, max }) => { + Self::merge_binary(op, Range::Scope { min, max }, Range::SortedRanges(ranges)) + .map_err(|(right, left)| (left, right)) } // e.g. c1 = 1 ? c1 = 2 (Range::Eq(left_val), Range::Eq(right_val)) => { - if left_val.eq(&right_val) && matches!(op, BinaryOperator::And | BinaryOperator::Or) - { - return Some(Range::Eq(left_val)); + if left_val == right_val && matches!(op, BinaryOperator::And | BinaryOperator::Or) { + return Ok(Range::Eq(left_val)); + } + if left_val.has_parameter() || right_val.has_parameter() { + return Err((Range::Eq(left_val), Range::Eq(right_val))); } match op { - BinaryOperator::And => Some(Range::Dummy), + BinaryOperator::And => Ok(Range::Dummy), BinaryOperator::Or => { - let mut ranges = Vec::new(); - let (val_1, val_2) = if let Some(true) = left_val.partial_cmp(&right_val).map(Ordering::is_gt) { @@ -548,31 +717,44 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { } else { (left_val, right_val) }; - ranges.push(Range::Eq(val_1)); - ranges.push(Range::Eq(val_2)); - Some(Range::SortedRanges(ranges)) + Ok(Range::SortedRanges(vec![ + Range::Eq(val_1), + Range::Eq(val_2), + ])) } - _ => None, + _ => Err((Range::Eq(left_val), Range::Eq(right_val))), } } // e.g. c1 = 1 ? (c1 = 1 or c1 = 2) - (Range::Eq(eq), Range::SortedRanges(ranges)) - | (Range::SortedRanges(ranges), Range::Eq(eq)) => { + (Range::Eq(eq), Range::SortedRanges(ranges)) => { + if eq.has_parameter() || ranges.iter().any(Range::has_parameter) { + return Err((Range::Eq(eq), Range::SortedRanges(ranges))); + } let merged_ranges = Self::extract_merge_ranges(op, Some(Range::Eq(eq)), ranges, &mut 0); - - Some(Self::ranges2range(merged_ranges)) + Ok(Self::ranges2range(merged_ranges)) + } + // e.g. (c1 = 1 or c1 = 2) ? c1 = 1 + (Range::SortedRanges(ranges), Range::Eq(eq)) => { + Self::merge_binary(op, Range::Eq(eq), Range::SortedRanges(ranges)) + .map_err(|(right, left)| (left, right)) } // e.g. (c1 = 1 or c1 = 2) ? (c1 = 1 or c1 = 2) (Range::SortedRanges(left_ranges), Range::SortedRanges(mut right_ranges)) => { + if left_ranges.iter().any(Range::has_parameter) + || right_ranges.iter().any(Range::has_parameter) + { + return Err(( + Range::SortedRanges(left_ranges), + Range::SortedRanges(right_ranges), + )); + } let mut idx = 0; - for left_range in left_ranges { right_ranges = Self::extract_merge_ranges(op, Some(left_range), right_ranges, &mut idx) } - - Some(Self::ranges2range(right_ranges)) + Ok(Self::ranges2range(right_ranges)) } } } @@ -608,17 +790,17 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { }, ) => { if let Some(true) = - Self::bound_compared(l_max, r_min, false).map(Ordering::is_lt) + Self::bound_compared(l_max, r_min, true, false).map(Ordering::is_lt) { ranges.insert(*idx, binary.unwrap()); return ranges; } else if let Some(true) = - Self::bound_compared(l_min, r_max, true).map(Ordering::is_gt) + Self::bound_compared(l_min, r_max, false, true).map(Ordering::is_gt) { *idx += 1; continue; } else { - binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)); + binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)).ok(); } } ( @@ -631,11 +813,11 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { let r_bound = Bound::Included(r_val.clone()); if let Some(true) = - Self::bound_compared(l_max, &r_bound, false).map(Ordering::is_lt) + Self::bound_compared(l_max, &r_bound, true, false).map(Ordering::is_lt) { ranges.insert(*idx, binary.unwrap()); return ranges; - } else if Self::bound_compared(l_min, &r_bound, true) + } else if Self::bound_compared(l_min, &r_bound, false, true) .map(Ordering::is_gt) .unwrap_or_else(|| op == BinaryOperator::Or) { @@ -644,7 +826,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { } else if r_val.is_null() { let _ = ranges.remove(*idx); } else { - binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)); + binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)).ok(); } } (Some(Range::Eq(l_val)), Range::Eq(r_val)) => { @@ -655,7 +837,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { *idx += 1; continue; } else { - binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)); + binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)).ok(); } } ( @@ -667,21 +849,21 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { ) => { let l_bound = Bound::Included(l_val.clone()); - if Self::bound_compared(&l_bound, r_min, false) + if Self::bound_compared(&l_bound, r_min, true, false) .map(Ordering::is_lt) .unwrap_or_else(|| op == BinaryOperator::Or) { ranges.insert(*idx, binary.unwrap()); return ranges; } else if let Some(true) = - Self::bound_compared(&l_bound, r_max, true).map(Ordering::is_gt) + Self::bound_compared(&l_bound, r_max, false, true).map(Ordering::is_gt) { *idx += 1; continue; } else if l_val.is_null() { binary = Some(ranges.remove(*idx)); } else { - binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)); + binary = Self::merge_binary(op, binary.unwrap(), ranges.remove(*idx)).ok(); } } (Some(Range::Dummy), _) => { @@ -722,14 +904,14 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { right_max: Bound, ) -> Range { if matches!( - Self::bound_compared(&left_max, &right_min, false), + Self::bound_compared(&left_max, &right_min, true, false), Some(Ordering::Less) ) || matches!( - Self::bound_compared(&right_max, &left_min, false), + Self::bound_compared(&right_max, &left_min, true, false), Some(Ordering::Less) ) { let (min_1, max_1, min_2, max_2) = if let Some(true) = - Self::bound_compared(&left_min, &right_min, true).map(Ordering::is_lt) + Self::bound_compared(&left_min, &right_min, false, false).map(Ordering::is_lt) { (left_min, left_max, right_min, right_max) } else { @@ -747,20 +929,20 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { ]); } let min = if let Some(true) = - Self::bound_compared(&left_min, &right_min, true).map(Ordering::is_lt) + Self::bound_compared(&left_min, &right_min, false, false).map(Ordering::is_lt) { left_min } else { right_min }; let max = if let Some(true) = - Self::bound_compared(&left_max, &right_max, false).map(Ordering::is_gt) + Self::bound_compared(&left_max, &right_max, true, true).map(Ordering::is_gt) { left_max } else { right_max }; - match Self::bound_compared(&min, &max, matches!(min, Bound::Unbounded)) { + match Self::bound_compared(&min, &max, false, true) { Some(Ordering::Equal) => match min { Bound::Included(val) => Range::Eq(val), Bound::Excluded(_) => Range::Dummy, @@ -778,22 +960,30 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { left_max: Bound, right_min: Bound, right_max: Bound, - ) -> Range { - let min = if let Some(true) = - Self::bound_compared(&left_min, &right_min, true).map(Ordering::is_gt) - { - left_min - } else { - right_min + ) -> Result { + let min_order = Self::bound_compared(&left_min, &right_min, false, false); + let max_order = Self::bound_compared(&left_max, &right_max, true, true); + let (Some(min_order), Some(max_order)) = (min_order, max_order) else { + return Err(( + Range::Scope { + min: left_min, + max: left_max, + }, + Range::Scope { + min: right_min, + max: right_max, + }, + )); }; - let max = if let Some(true) = - Self::bound_compared(&left_max, &right_max, false).map(Ordering::is_lt) - { - left_max - } else { - right_max + let min = match min_order { + Ordering::Greater => left_min, + Ordering::Less | Ordering::Equal => right_min, }; - match Self::bound_compared(&min, &max, matches!(min, Bound::Unbounded)) { + let max = match max_order { + Ordering::Less => left_max, + Ordering::Greater | Ordering::Equal => right_max, + }; + Ok(match Self::bound_compared(&min, &max, false, true) { Some(Ordering::Greater) => Range::Dummy, Some(Ordering::Equal) => match min { Bound::Included(val) => Range::Eq(val), @@ -804,7 +994,7 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { }, }, _ => Range::Scope { min, max }, - } + }) } fn matches_column(&self, col: ColumnRef) -> bool { @@ -820,33 +1010,56 @@ impl<'a, 'p, M: RangeColumnMatcher> RangeDetacher<'a, 'p, M> { fn bound_compared( left_bound: &Bound, right_bound: &Bound, - is_min: bool, + left_is_upper: bool, + right_is_upper: bool, ) -> Option { - fn is_min_then_reverse(is_min: bool, order: Ordering) -> Ordering { - if is_min { - order - } else { - order.reverse() - } - } - fn range_value_cmp(left: &DataValue, right: &DataValue) -> Option { + fn range_value_cmp( + left: &DataValue, + right: &DataValue, + left_is_upper: bool, + right_is_upper: bool, + ) -> Option { match (left, right) { (DataValue::Null, DataValue::Null) => Some(Ordering::Equal), (DataValue::Null, _) => Some(Ordering::Greater), (_, DataValue::Null) => Some(Ordering::Less), + (DataValue::Tuple(left), DataValue::Tuple(right)) => { + crate::types::value::tuple_partial_cmp( + left, + right, + left_is_upper, + right_is_upper, + ) + } _ => left.partial_cmp(right), } } + fn is_min_then_reverse(is_min: bool, order: Ordering) -> Ordering { + if is_min { + order + } else { + order.reverse() + } + } + let is_min = !left_is_upper && !right_is_upper; match (left_bound, right_bound) { (Bound::Unbounded, Bound::Unbounded) => Some(Ordering::Equal), - (Bound::Unbounded, _) => Some(is_min_then_reverse(is_min, Ordering::Less)), - (_, Bound::Unbounded) => Some(is_min_then_reverse(is_min, Ordering::Greater)), - (Bound::Included(left), Bound::Included(right)) => range_value_cmp(left, right), - (Bound::Included(left), Bound::Excluded(right)) => range_value_cmp(left, right) - .map(|order| order.then(is_min_then_reverse(is_min, Ordering::Less))), - (Bound::Excluded(left), Bound::Excluded(right)) => range_value_cmp(left, right), - (Bound::Excluded(left), Bound::Included(right)) => range_value_cmp(left, right) - .map(|order| order.then(is_min_then_reverse(is_min, Ordering::Greater))), + (Bound::Unbounded, _) => Some(is_min_then_reverse(!left_is_upper, Ordering::Less)), + (_, Bound::Unbounded) => Some(is_min_then_reverse(!right_is_upper, Ordering::Greater)), + (Bound::Included(left), Bound::Included(right)) => { + range_value_cmp(left, right, left_is_upper, right_is_upper) + } + (Bound::Included(left), Bound::Excluded(right)) => { + range_value_cmp(left, right, left_is_upper, !right_is_upper) + .map(|order| order.then(is_min_then_reverse(is_min, Ordering::Less))) + } + (Bound::Excluded(left), Bound::Excluded(right)) => { + range_value_cmp(left, right, !left_is_upper, !right_is_upper) + } + (Bound::Excluded(left), Bound::Included(right)) => { + range_value_cmp(left, right, !left_is_upper, right_is_upper) + .map(|order| order.then(is_min_then_reverse(is_min, Ordering::Greater))) + } } } @@ -954,17 +1167,18 @@ pub(crate) mod test_support { } impl RangeColumnMatcher for DirectRangeColumn<'_> { - fn matches(&self, table_name: &str, column_id: ColumnId, _arena: &PlanArena<'_>) -> bool { + fn matches( + &self, + table_name: &str, + column_id: ColumnId, + _arena: &A, + ) -> 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 { + impl<'a, A: MetaArena + ?Sized> RangeDetacher<'a, DirectRangeColumn<'a>, A> { + pub(crate) fn new(table_name: &'a str, column_id: &'a ColumnId, arena: &'a mut A) -> Self { Self { column: DirectRangeColumn { table_name, @@ -989,15 +1203,14 @@ mod test { use crate::optimizer::rule::normalization::NormalizationRuleImpl; use crate::planner::operator::filter::FilterOperator; use crate::planner::operator::Operator; - use crate::planner::{ExprRef, LogicalPlan}; - use crate::types::evaluator::binary_create; + use crate::planner::{ExprRef, LogicalPlan, MetaArena, PlanArena}; use crate::types::value::DataValue; use crate::types::LogicalType; use std::ops::Bound; fn plan_filter( plan: LogicalPlan, - arena: &mut crate::planner::PlanArena, + arena: &mut PlanArena<'_>, ) -> Result, DatabaseError> { let pipeline = HepOptimizerPipeline::builder() .before_batch( @@ -1015,7 +1228,7 @@ mod test { } fn test_column( - arena: &mut crate::planner::PlanArena, + arena: &mut (dyn MetaArena + '_), table_name: &TableName, column_id: crate::types::ColumnId, name: &str, @@ -1031,7 +1244,7 @@ mod test { } fn cmp_predicate( - arena: &mut crate::planner::PlanArena, + arena: &mut (dyn MetaArena + '_), column: ColumnRef, op: BinaryOperator, value: i32, @@ -1048,11 +1261,7 @@ mod test { }) } - fn and_predicate( - arena: &mut crate::planner::PlanArena, - left: ExprRef, - right: ExprRef, - ) -> ExprRef { + fn and_predicate(arena: &mut (dyn MetaArena + '_), left: ExprRef, right: ExprRef) -> ExprRef { arena.alloc_expression(ScalarExpression::Binary { op: BinaryOperator::And, left_expr: left, @@ -2147,110 +2356,82 @@ mod test { range, Some(Range::SortedRanges(vec![ Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![DataValue::Int32(1), DataValue::Null, DataValue::Int32(1),], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![DataValue::Int32(1), DataValue::Null, DataValue::Int32(2),], - true - )) + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + ])) }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + ])), }, ])) ); @@ -2265,110 +2446,82 @@ mod test { range, Some(Range::SortedRanges(vec![ Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![DataValue::Int32(1), DataValue::Null, DataValue::Int32(1),], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(1), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![DataValue::Int32(1), DataValue::Null, DataValue::Int32(2),], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(2), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + DataValue::Int32(1), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - DataValue::Int32(1), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + DataValue::Int32(1), + ])), }, ])) ); @@ -2383,124 +2536,88 @@ mod test { range, Some(Range::SortedRanges(vec![ Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(1), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(1), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Null, - DataValue::Int32(2), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Null, + DataValue::Int32(2), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(1), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(1), + DataValue::Int32(2), + ])), }, Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - DataValue::Int32(1), - ], - false - )), - max: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(2), - DataValue::Int32(2), - ], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + DataValue::Int32(1), + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(2), + DataValue::Int32(2), + ])), }, ])) ) @@ -2519,35 +2636,17 @@ mod test { assert_eq!( range, Some(Range::Scope { - min: Bound::Included(DataValue::Tuple( - vec![ - DataValue::Int32(7), - DataValue::Int32(10), - DataValue::Int32(2) - ], - false - )), - max: Bound::Excluded(DataValue::Tuple( - vec![DataValue::Int32(7), DataValue::Int32(10)], - true - )), + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(7), + DataValue::Int32(10), + DataValue::Int32(2) + ])), + max: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(7), + DataValue::Int32(10) + ])), }) ); - let Range::Scope { - min: Bound::Included(min), - max: Bound::Excluded(max), - } = range.unwrap() - else { - unreachable!() - }; - assert_eq!( - binary_create( - std::borrow::Cow::Owned(LogicalType::Tuple(vec![])), - BinaryOperator::Lt - )? - .binary_eval(&min, &max)?, - DataValue::Boolean(true) - ); Ok(()) } @@ -2563,7 +2662,7 @@ mod test { gt_one.clone(), Range::Eq(DataValue::Int32(1)), ), - Some(Range::Scope { + Ok(Range::Scope { min: Bound::Included(DataValue::Int32(1)), max: Bound::Unbounded, }) @@ -2574,7 +2673,7 @@ mod test { gt_one, Range::Eq(DataValue::Int32(1)), ), - Some(Range::Dummy) + Ok(Range::Dummy) ); let disjoint = RangeDetacher::::merge_binary( @@ -2590,7 +2689,7 @@ mod test { ); assert_eq!( disjoint, - Some(Range::SortedRanges(vec![ + Ok(Range::SortedRanges(vec![ Range::Scope { min: Bound::Included(DataValue::Int32(1)), max: Bound::Included(DataValue::Int32(2)), @@ -2601,6 +2700,23 @@ mod test { }, ])) ); + + let left = Range::Eq(DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + }); + let right = Range::Eq(DataValue::Parameter { + id: 2, + ty: LogicalType::Integer, + }); + assert_eq!( + RangeDetacher::::merge_binary( + BinaryOperator::And, + left.clone(), + right.clone(), + ), + Err((left, right)) + ); } #[test] @@ -2621,6 +2737,57 @@ mod test { assert!(prefixes[0].only_eq()); assert!(prefixes[1].only_eq()); + let explicit_suffix = Range::Scope { + min: Bound::Included(DataValue::Int32(10)), + max: Bound::Excluded(DataValue::Int32(20)), + } + .combining_eqs(&[Range::Eq(DataValue::Int32(1))]); + assert_eq!( + explicit_suffix, + Some(Range::Scope { + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(10) + ],)), + max: Bound::Excluded(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(20) + ],)), + }) + ); + + let unbounded_suffix = Range::Scope { + min: Bound::Unbounded, + max: Bound::Unbounded, + } + .combining_eqs(&[Range::Eq(DataValue::Int32(1))]); + assert_eq!( + unbounded_suffix, + Some(Range::Scope { + min: Bound::Included(DataValue::Tuple(vec![DataValue::Int32(1)])), + max: Bound::Included(DataValue::Tuple(vec![DataValue::Int32(1)])), + }) + ); + + let concrete = Range::Scope { + min: Bound::Included(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(10), + ])), + max: Bound::Excluded(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(20), + ])), + }; + assert_eq!( + RangeDetacher::::merge_binary( + BinaryOperator::And, + unbounded_suffix.unwrap(), + concrete.clone(), + ), + Ok(concrete) + ); + assert!(suffix .combining_eqs(&[Range::Scope { min: Bound::Unbounded, diff --git a/src/expression/simplify.rs b/src/expression/simplify.rs index 5379e65d..8823b898 100644 --- a/src/expression/simplify.rs +++ b/src/expression/simplify.rs @@ -16,7 +16,8 @@ use crate::catalog::ColumnRef; use crate::errors::DatabaseError; use crate::expression::visitor_mut::ExprVisitorMut; use crate::expression::{BinaryOperator, ScalarExpression, TypeCast, UnaryOperator}; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::evaluator::{binary_create, unary_create}; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -47,7 +48,7 @@ struct ReplaceUnary { pub struct ConstantCalculator; impl ConstantCalculator { - pub fn new(_arena: &PlanArena<'_>) -> Self { + pub fn new(_arena: &(dyn MetaArena + '_)) -> Self { Self } } @@ -56,7 +57,7 @@ impl ExprVisitorMut for ConstantCalculator { fn visit_expression( &mut self, expr: &mut ScalarExpression, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result { match expr { ScalarExpression::Unary { @@ -68,6 +69,9 @@ impl ExprVisitorMut for ConstantCalculator { self.visit(arg_expr, arena)?; if let ScalarExpression::Constant(unary_val) = arena.expression(*arg_expr) { + if unary_val.has_parameter() { + return Ok(false); + } let value = if let Some(evaluator) = evaluator { evaluator.unary_eval(unary_val) } else { @@ -95,6 +99,9 @@ impl ExprVisitorMut for ConstantCalculator { ScalarExpression::Constant(right_val), ) = (arena.expression(*left_expr), arena.expression(*right_expr)) { + if left_val.has_parameter() || right_val.has_parameter() { + return Ok(false); + } let evaluator = binary_create(Cow::Borrowed(&ty), *op)?; let left_val = left_val.clone().cast(&ty)?; let right_val = right_val.clone().cast(&ty)?; @@ -108,6 +115,9 @@ impl ExprVisitorMut for ConstantCalculator { self.visit(arg_expr, arena)?; if let ScalarExpression::Constant(value) = arena.expression(*arg_expr) { + if value.has_parameter() { + return Ok(false); + } let casted = value.clone().cast(ty)?; *expr = ScalarExpression::Constant(casted); } @@ -128,7 +138,7 @@ impl ExprVisitorMut for Simplify { fn visit_expression( &mut self, expr: &mut ScalarExpression, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result { match expr { ScalarExpression::Unary { @@ -351,6 +361,18 @@ impl Simplify { ) } + fn is_rearrangeable_comparison(op: &BinaryOperator) -> bool { + matches!( + op, + BinaryOperator::Gt + | BinaryOperator::Lt + | BinaryOperator::GtEq + | BinaryOperator::LtEq + | BinaryOperator::Eq + | BinaryOperator::NotEq + ) + } + fn negate_range_comparison(op: BinaryOperator) -> Option { match op { BinaryOperator::Gt => Some(BinaryOperator::LtEq), @@ -361,7 +383,10 @@ impl Simplify { } } - fn take_range_comparison(expr: ExprRef, arena: &PlanArena<'_>) -> Option { + fn take_range_comparison( + expr: ExprRef, + arena: &(dyn MetaArena + '_), + ) -> Option { match arena.expression(expr) { expression @ ScalarExpression::Binary { op, .. } if Self::negate_range_comparison(*op).is_some() => @@ -374,7 +399,7 @@ impl Simplify { fn take_negated_range_comparison( expr: ExprRef, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Option { let mut expression = arena.expression(expr).clone(); match &mut expression { @@ -386,7 +411,7 @@ impl Simplify { } } - fn boolean_constant(expr: ExprRef, arena: &PlanArena<'_>) -> Option { + fn boolean_constant(expr: ExprRef, arena: &(dyn MetaArena + '_)) -> Option { match arena.expression(expr) { ScalarExpression::Constant(DataValue::Boolean(value)) => Some(*value), _ => None, @@ -396,7 +421,7 @@ impl Simplify { fn take_range_comparison_with_polarity( expr: ExprRef, positive: bool, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Option { if positive { Self::take_range_comparison(expr, arena) @@ -409,7 +434,7 @@ impl Simplify { op: BinaryOperator, left_expr: ExprRef, right_expr: ExprRef, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Option { let is_eq = matches!(op, BinaryOperator::Eq); if !matches!(op, BinaryOperator::Eq | BinaryOperator::NotEq) { @@ -439,13 +464,19 @@ impl Simplify { left_expr: &mut ExprRef, right_expr: &mut ExprRef, op: &mut BinaryOperator, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(left_expr, arena)?; if Self::is_arithmetic(op) { return Ok(()); } + // Terms can only be moved across a comparison. Operators such as `%` + // are not invertible, so pending replaces must not be applied to them. + if !Self::is_rearrangeable_comparison(op) { + self.replaces.clear(); + return Ok(()); + } while let Some(replace) = self.replaces.pop() { match replace { Replace::Binary(binary) => { @@ -466,7 +497,7 @@ impl Simplify { col_expr: &mut ExprRef, val_expr: &mut ExprRef, op: &mut BinaryOperator, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) { let ReplaceUnary { child_expr, @@ -509,7 +540,7 @@ impl Simplify { left_expr: &mut ExprRef, right_expr: &mut ExprRef, op: &mut BinaryOperator, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) { let ReplaceBinary { column_expr, @@ -553,28 +584,35 @@ impl Simplify { } impl ExprRef { - pub(crate) fn unpack_val(self, arena: &PlanArena<'_>) -> Option { + pub(crate) fn unpack_val(self, arena: &A) -> Option { match arena.expression(self) { ScalarExpression::Constant(val) => Some(val.clone()), ScalarExpression::Alias { expr, .. } => expr.unpack_val(arena), ScalarExpression::TypeCast { expr, ty, .. } => { expr.unpack_val(arena).and_then(|val| val.cast(ty).ok()) } - ScalarExpression::IsNull { negated, expr } => Some(DataValue::Boolean( - expr.unpack_val(arena)?.is_null() != *negated, - )), + ScalarExpression::IsNull { negated, expr } => { + let value = expr.unpack_val(arena)?; + (!value.has_parameter()).then(|| DataValue::Boolean(value.is_null() != *negated)) + } ScalarExpression::Unary { expr, op, evaluator, ty, - } => 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)?) - }), + } => { + let value = expr.unpack_val(arena)?; + if value.has_parameter() { + return None; + } + Some(if let Some(evaluator) = evaluator { + evaluator.unary_eval(&value) + } else { + unary_create(Cow::Borrowed(ty), *op) + .ok()? + .unary_eval(&value) + }) + } ScalarExpression::Binary { left_expr, right_expr, @@ -584,6 +622,9 @@ impl ExprRef { } => { let left = left_expr.unpack_val(arena)?.cast(ty).ok()?; let right = right_expr.unpack_val(arena)?.cast(ty).ok()?; + if left.has_parameter() || right.has_parameter() { + return None; + } if let Some(evaluator) = evaluator { evaluator.binary_eval(&left, &right) } else { @@ -597,9 +638,9 @@ impl ExprRef { } } - pub(crate) fn unpack_bound_col( + pub(crate) fn unpack_bound_col( self, - arena: &PlanArena<'_>, + arena: &A, is_deep: bool, ) -> Option<(ColumnRef, usize)> { match arena.expression(self) { diff --git a/src/expression/visitor.rs b/src/expression/visitor.rs index e1629623..a4b653f2 100644 --- a/src/expression/visitor.rs +++ b/src/expression/visitor.rs @@ -26,7 +26,7 @@ use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluat use crate::types::value::DataValue; use crate::types::LogicalType; -pub trait ExprVisitor: Sized { +pub trait ExprVisitor: Sized { fn visit(&mut self, expr: ExprRef, arena: &A) -> Result<(), DatabaseError> { if !self.visit_expression_ref(expr, arena)? { return Ok(()); @@ -287,7 +287,7 @@ pub trait ExprVisitor: Sized { } } -pub fn walk_expr>( +pub fn walk_expr>( visitor: &mut V, expr: ExprRef, arena: &A, diff --git a/src/expression/visitor_mut.rs b/src/expression/visitor_mut.rs index 4eea5c5a..f9c1ad7c 100644 --- a/src/expression/visitor_mut.rs +++ b/src/expression/visitor_mut.rs @@ -21,7 +21,8 @@ use crate::expression::window::WindowCall; use crate::expression::{ AliasType, BinaryOperator, ScalarExpression, TrimWhereField, UnaryOperator, }; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::evaluator::{BinaryEvaluatorRef, CastEvaluatorRef, UnaryEvaluatorRef}; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -32,7 +33,7 @@ impl ExprVisitorMut for ExprCloner { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { *expr = arena.alloc_expression(arena.expression(*expr).clone()); walk_mut_expr(self, expr, arena) @@ -48,7 +49,7 @@ impl ExprVisitorMut for PositionShift { &mut self, _column: &mut ColumnRef, position: &mut usize, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if self.delta.is_negative() { *position = position.saturating_sub(self.delta.unsigned_abs()); @@ -63,14 +64,14 @@ pub trait ExprVisitorMut: Sized { fn visit( &mut self, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if !self.visit_expression_ref(expr, arena)? { return Ok(()); } let mut expression = - std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + std::mem::replace(&mut *arena.expression_mut(*expr), ScalarExpression::Empty); let result = self.visit_expression(&mut expression, arena); *arena.expression_mut(*expr) = expression; if result? { @@ -82,7 +83,7 @@ pub trait ExprVisitorMut: Sized { fn visit_expression_ref( &mut self, _expr: &mut ExprRef, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result { Ok(true) } @@ -90,7 +91,7 @@ pub trait ExprVisitorMut: Sized { fn visit_expression( &mut self, _expr: &mut ScalarExpression, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result { Ok(true) } @@ -98,7 +99,7 @@ pub trait ExprVisitorMut: Sized { fn visit_constant( &mut self, _value: &mut DataValue, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { Ok(()) } @@ -107,7 +108,7 @@ pub trait ExprVisitorMut: Sized { &mut self, _column: &mut ColumnRef, _position: &mut usize, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { Ok(()) } @@ -116,7 +117,7 @@ pub trait ExprVisitorMut: Sized { &mut self, expr: &mut ExprRef, alias: &mut AliasType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let AliasType::Expr(alias_expr) = alias { self.visit(alias_expr, arena)?; @@ -129,7 +130,7 @@ pub trait ExprVisitorMut: Sized { expr: &mut ExprRef, _ty: &mut LogicalType, _evaluator: &mut Option, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena) } @@ -138,7 +139,7 @@ pub trait ExprVisitorMut: Sized { &mut self, _negated: bool, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena) } @@ -149,7 +150,7 @@ pub trait ExprVisitorMut: Sized { expr: &mut ExprRef, _evaluator: &mut Option, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena) } @@ -161,7 +162,7 @@ pub trait ExprVisitorMut: Sized { right_expr: &mut ExprRef, _evaluator: &mut Option, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(left_expr, arena)?; self.visit(right_expr, arena) @@ -173,7 +174,7 @@ pub trait ExprVisitorMut: Sized { _kind: &mut AggKind, args: &mut [ExprRef], _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for arg in args { self.visit(arg, arena)?; @@ -184,7 +185,7 @@ pub trait ExprVisitorMut: Sized { fn visit_window( &mut self, window: &mut WindowCall, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for expr in window .function @@ -203,7 +204,7 @@ pub trait ExprVisitorMut: Sized { _negated: bool, expr: &mut ExprRef, args: &mut [ExprRef], - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; for arg in args { @@ -218,7 +219,7 @@ pub trait ExprVisitorMut: Sized { expr: &mut ExprRef, left_expr: &mut ExprRef, right_expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; self.visit(left_expr, arena)?; @@ -230,7 +231,7 @@ pub trait ExprVisitorMut: Sized { expr: &mut ExprRef, for_expr: &mut Option, from_expr: &mut Option, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; if let Some(for_expr) = for_expr { @@ -246,7 +247,7 @@ pub trait ExprVisitorMut: Sized { &mut self, expr: &mut ExprRef, in_expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; self.visit(in_expr, arena) @@ -257,7 +258,7 @@ pub trait ExprVisitorMut: Sized { expr: &mut ExprRef, trim_what_expr: &mut Option, _trim_where: &mut Option, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena)?; if let Some(trim_what_expr) = trim_what_expr { @@ -274,7 +275,7 @@ pub trait ExprVisitorMut: Sized { &mut self, expr: &mut ExprRef, _pos: usize, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena) } @@ -282,7 +283,7 @@ pub trait ExprVisitorMut: Sized { fn visit_tuple( &mut self, exprs: &mut [ExprRef], - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for expr in exprs { self.visit(expr, arena)?; @@ -293,7 +294,7 @@ pub trait ExprVisitorMut: Sized { fn visit_scala_function( &mut self, function: &mut ScalarFunction, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for arg in &mut function.args { self.visit(arg, arena)?; @@ -304,7 +305,7 @@ pub trait ExprVisitorMut: Sized { fn visit_table_function( &mut self, function: &mut TableFunction, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for arg in &mut function.args { self.visit(arg, arena)?; @@ -318,7 +319,7 @@ pub trait ExprVisitorMut: Sized { left_expr: &mut ExprRef, right_expr: &mut ExprRef, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(condition, arena)?; self.visit(left_expr, arena)?; @@ -330,7 +331,7 @@ pub trait ExprVisitorMut: Sized { left_expr: &mut ExprRef, right_expr: &mut ExprRef, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(left_expr, arena)?; self.visit(right_expr, arena) @@ -341,7 +342,7 @@ pub trait ExprVisitorMut: Sized { left_expr: &mut ExprRef, right_expr: &mut ExprRef, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(left_expr, arena)?; self.visit(right_expr, arena) @@ -351,7 +352,7 @@ pub trait ExprVisitorMut: Sized { &mut self, exprs: &mut [ExprRef], _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { for expr in exprs { self.visit(expr, arena)?; @@ -365,7 +366,7 @@ pub trait ExprVisitorMut: Sized { expr_pairs: &mut [(ExprRef, ExprRef)], else_expr: &mut Option, _ty: &mut LogicalType, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let Some(expr) = operand_expr { self.visit(expr, arena)?; @@ -384,9 +385,10 @@ pub trait ExprVisitorMut: Sized { pub fn walk_mut_expr( visitor: &mut V, expr: &mut ExprRef, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { - let mut expression = std::mem::replace(arena.expression_mut(*expr), ScalarExpression::Empty); + let mut expression = + std::mem::replace(&mut *arena.expression_mut(*expr), ScalarExpression::Empty); let result = match &mut expression { ScalarExpression::Constant(value) => visitor.visit_constant(value, arena), ScalarExpression::ColumnRef { column, position } => { diff --git a/src/function/char_length.rs b/src/function/char_length.rs index 73352dba..8a413802 100644 --- a/src/function/char_length.rs +++ b/src/function/char_length.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -44,15 +45,19 @@ impl ScalarFunctionImpl for CharLength { fn eval( &self, exprs: &[ExprRef], - arena: &crate::planner::PlanArena<'_>, + arena: &(dyn MetaArena + '_), tuples: Option<&dyn TupleLike>, ) -> Result { 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))?; + value = std::borrow::Cow::Owned( + value + .into_owned() + .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?, + ); } let mut length: u64 = 0; - if let DataValue::Utf8 { value, ty, unit } = &mut value { + if let DataValue::Utf8 { value, ty, unit } = value.as_ref() { length = value.chars().count() as u64; } Ok(DataValue::UInt64(length)) diff --git a/src/function/current_date.rs b/src/function/current_date.rs index f3a5ea30..1deb47bf 100644 --- a/src/function/current_date.rs +++ b/src/function/current_date.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -46,7 +47,7 @@ impl ScalarFunctionImpl for CurrentDate { fn eval( &self, _: &[ExprRef], - _: &crate::planner::PlanArena<'_>, + _: &(dyn MetaArena + '_), _: 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 7c443e8c..b8a00f4d 100644 --- a/src/function/current_timestamp.rs +++ b/src/function/current_timestamp.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::LogicalType; @@ -46,7 +47,7 @@ impl ScalarFunctionImpl for CurrentTimeStamp { fn eval( &self, _: &[ExprRef], - _: &crate::planner::PlanArena<'_>, + _: &(dyn MetaArena + '_), _: 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 c2396395..8e4e8b53 100644 --- a/src/function/lower.rs +++ b/src/function/lower.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -46,17 +47,25 @@ impl ScalarFunctionImpl for Lower { fn eval( &self, exprs: &[ExprRef], - arena: &crate::planner::PlanArena<'_>, + arena: &(dyn MetaArena + '_), tuples: Option<&dyn TupleLike>, ) -> Result { 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))?; + value = std::borrow::Cow::Owned( + value + .into_owned() + .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?, + ); } - if let DataValue::Utf8 { value, ty, unit } = &mut value { - *value = value.to_lowercase(); - } - Ok(value) + Ok(match value.as_ref() { + DataValue::Utf8 { value, ty, unit } => DataValue::Utf8 { + value: value.to_lowercase(), + ty: ty.clone(), + unit: *unit, + }, + _ => value.into_owned(), + }) } fn monotonicity(&self) -> Option { diff --git a/src/function/numbers.rs b/src/function/numbers.rs index 08716390..f5c844e8 100644 --- a/src/function/numbers.rs +++ b/src/function/numbers.rs @@ -17,6 +17,7 @@ use crate::catalog::ColumnDesc; use crate::errors::DatabaseError; use crate::expression::function::table::TableFunctionImpl; use crate::expression::function::FunctionSummary; +use crate::planner::MetaArena; use crate::planner::{ExprRef, TableArena}; use crate::types::tuple::Schema; use crate::types::tuple::Tuple; @@ -47,9 +48,9 @@ impl TableFunctionImpl for Numbers { fn eval( &self, args: &[ExprRef], - arena: &crate::planner::PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result>>, DatabaseError> { - let mut value = arena.expression(args[0]).eval::<&Tuple>(arena, None)?; + let mut value = arena.expression(args[0]).eval(arena, None)?.into_owned(); value = value.cast(&LogicalType::Integer)?; let num = value diff --git a/src/function/octet_length.rs b/src/function/octet_length.rs index 4c31d80d..273ff930 100644 --- a/src/function/octet_length.rs +++ b/src/function/octet_length.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -45,15 +46,19 @@ impl ScalarFunctionImpl for OctetLength { fn eval( &self, exprs: &[ExprRef], - arena: &crate::planner::PlanArena<'_>, + arena: &(dyn MetaArena + '_), tuples: Option<&dyn TupleLike>, ) -> Result { 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))?; + value = std::borrow::Cow::Owned( + value + .into_owned() + .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?, + ); } let mut length: u64 = 0; - if let DataValue::Utf8 { value, ty, unit } = &mut value { + if let DataValue::Utf8 { value, ty, unit } = value.as_ref() { length = value.len() as u64; } Ok(DataValue::UInt64(length)) diff --git a/src/function/upper.rs b/src/function/upper.rs index f2a8ac1c..c6da1e46 100644 --- a/src/function/upper.rs +++ b/src/function/upper.rs @@ -17,6 +17,7 @@ use crate::expression::function::scala::FuncMonotonicity; use crate::expression::function::scala::ScalarFunctionImpl; use crate::expression::function::FunctionSummary; use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::tuple::TupleLike; use crate::types::value::DataValue; use crate::types::CharLengthUnits; @@ -46,17 +47,25 @@ impl ScalarFunctionImpl for Upper { fn eval( &self, exprs: &[ExprRef], - arena: &crate::planner::PlanArena<'_>, + arena: &(dyn MetaArena + '_), tuples: Option<&dyn TupleLike>, ) -> Result { 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))?; + value = std::borrow::Cow::Owned( + value + .into_owned() + .cast(&LogicalType::Varchar(None, CharLengthUnits::Characters))?, + ); } - if let DataValue::Utf8 { value, ty, unit } = &mut value { - *value = value.to_uppercase(); - } - Ok(value) + Ok(match value.as_ref() { + DataValue::Utf8 { value, ty, unit } => DataValue::Utf8 { + value: value.to_uppercase(), + ty: ty.clone(), + unit: *unit, + }, + _ => value.into_owned(), + }) } fn monotonicity(&self) -> Option { diff --git a/src/macros/mod.rs b/src/macros/mod.rs index e3af1a72..9db99307 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::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> { + fn eval(&self, args: &[::kite_sql::planner::ExprRef], arena: &dyn ::kite_sql::planner::MetaArena, 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 = arena.expression(args[_index]).eval(arena, tuple)?; + let mut value = arena.expression(args[_index]).eval(arena, tuple)?.into_owned(); _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::planner::ExprRef], arena: &::kite_sql::planner::PlanArena<'_>) -> Result>>, ::kite_sql::errors::DatabaseError> { + fn eval(&self, args: &[::kite_sql::planner::ExprRef], arena: &dyn ::kite_sql::planner::MetaArena) -> Result>>, ::kite_sql::errors::DatabaseError> { let mut _index = 0; $closure($({ - let mut value = arena.expression(args[_index]).eval::<&::kite_sql::types::tuple::Tuple>(arena, None)?; + let mut value = arena.expression(args[_index]).eval(arena, None)?.into_owned(); _index += 1; value = value.cast(&$arg_ty)?; diff --git a/src/optimizer/core/cm_sketch.rs b/src/optimizer/core/cm_sketch.rs index afe2d8d8..944ac77d 100644 --- a/src/optimizer/core/cm_sketch.rs +++ b/src/optimizer/core/cm_sketch.rs @@ -14,6 +14,7 @@ use crate::errors::DatabaseError; use crate::expression::range_detacher::Range; +use crate::planner::MetaArena; use crate::serdes::stable_hash::{StableHasher, CM_SKETCH_HASH_KEYS}; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; @@ -353,7 +354,7 @@ impl CountMinSketch { } impl ReferenceSerialization for CountMinSketch { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -374,7 +375,7 @@ impl ReferenceSerialization for CountMinSketch { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/optimizer/core/histogram.rs b/src/optimizer/core/histogram.rs index 343f423b..d1352a58 100644 --- a/src/optimizer/core/histogram.rs +++ b/src/optimizer/core/histogram.rs @@ -392,6 +392,31 @@ impl Histogram { self.equal_count(value, sketch).clamp(lower, entry.count()) } + fn estimate_range_selectivity(&self, range: &Range) -> f64 { + let Some(bucket) = self.buckets.first() else { + return 0.0; + }; + let key_type = bucket.upper.logical_type(); + let distinct = self.distinct_values_len().max(1) as f64; + match range { + Range::Dummy => 0.0, + Range::Eq(value) => DataValue::bound_selectivity( + Bound::Included(value), + Bound::Included(value), + &key_type, + distinct, + ), + Range::Scope { min, max } => { + DataValue::bound_selectivity(min.as_ref(), max.as_ref(), &key_type, distinct) + } + Range::SortedRanges(ranges) => ranges + .iter() + .map(|range| self.estimate_range_selectivity(range)) + .sum::() + .min(1.0), + } + } + pub fn collect_count( &self, ranges: &[Range], @@ -401,12 +426,26 @@ impl Histogram { if self.buckets.is_empty() || ranges.is_empty() { return Ok(0); } + if ranges.iter().any(Range::has_parameter) { + let selectivity = ranges + .iter() + .map(|range| self.estimate_range_selectivity(range)) + .sum::() + .min(1.0); + let count = (self.values_len() as f64 * selectivity).ceil() as usize; + return Ok(if selectivity > 0.0 && self.values_len() > 0 { + count.max(1) + } else { + 0 + }); + } let comparator = self.comparator()?; let mut count = 0; let mut binary_i = 0; let mut bucket_i = 0; let mut bucket_idxs = Vec::new(); + let mut buf = Vec::new(); while bucket_i < self.buckets.len() && binary_i < ranges.len() { let is_dummy = self._collect_count( @@ -418,6 +457,7 @@ impl Histogram { sketch, top_n, comparator, + &mut buf, )?; if is_dummy { return Ok(0); @@ -442,96 +482,8 @@ impl Histogram { sketch: &CountMinSketch, top_n: &ColumnTopN, comparator: &BoundComparator, + buf: &mut Vec, ) -> Result { - let float_value = |value: &DataValue, prefix_len: usize| { - let value = match value.logical_type() { - LogicalType::Varchar(..) | LogicalType::Char(..) => match value { - DataValue::Utf8 { value, .. } => { - if prefix_len > value.len() { - return Ok(0.0); - } - - let mut val = 0u64; - for (i, char) in value - .get(prefix_len..prefix_len + 8) - .unwrap() - .chars() - .enumerate() - { - if value.len() - prefix_len > i { - val += (val << 8) + char as u64; - } else { - val += val << 8; - } - } - - Some(val as f64) - } - _ => unreachable!(), - }, - LogicalType::Date - | LogicalType::DateTime - | LogicalType::Time(_) - | LogicalType::TimeStamp(_, _) => match value { - DataValue::Date32(value) => DataValue::Int32(*value) - .cast(&LogicalType::Double)? - .double(), - DataValue::Date64(value) => DataValue::Int64(*value) - .cast(&LogicalType::Double)? - .double(), - DataValue::Time32(value, ..) => DataValue::UInt32(*value) - .cast(&LogicalType::Double)? - .double(), - DataValue::Time64(value, ..) => DataValue::Int64(*value) - .cast(&LogicalType::Double)? - .double(), - _ => unreachable!(), - }, - - LogicalType::SqlNull - | LogicalType::Boolean - | LogicalType::Tinyint - | LogicalType::UTinyint - | LogicalType::Smallint - | LogicalType::USmallint - | LogicalType::Integer - | LogicalType::UInteger - | LogicalType::Bigint - | LogicalType::UBigint - | LogicalType::Float - | LogicalType::Double - | LogicalType::Decimal(_, _) => value.clone().cast(&LogicalType::Double)?.double(), - LogicalType::Tuple(_) => match value { - DataValue::Tuple(values, _) => { - let mut float = 0.0; - - for (i, value) in values.iter().enumerate() { - if !value.logical_type().is_numeric() { - continue; - } - if let Some(f) = - DataValue::clone(value).cast(&LogicalType::Double)?.double() - { - float += f / (10_i32.pow(i as u32) as f64); - } - } - Some(float) - } - DataValue::Null => None, - _ => unreachable!(), - }, - } - .unwrap_or(0.0); - Ok::(value) - }; - let calc_fraction = |start: &DataValue, end: &DataValue, value: &DataValue| { - let prefix_len = start.common_prefix_length(end).unwrap_or(0); - Ok::( - (float_value(value, prefix_len)? - float_value(start, prefix_len)?) - / (float_value(end, prefix_len)? - float_value(start, prefix_len)?), - ) - }; - let distinct_1 = OrderedFloat(1.0 / self.meta.number_of_distinct_value as f64); match &ranges[*binary_i] { @@ -560,11 +512,12 @@ impl Histogram { *bucket_i += 1; } else if is_above(comparator, &bucket.lower, min, true)? { let (temp_ratio, option) = match max { - Bound::Included(val) => { - (calc_fraction(&bucket.lower, &bucket.upper, val)?, None) - } + Bound::Included(val) => ( + encoded_fraction(&bucket.lower, &bucket.upper, val, true, buf)?, + None, + ), Bound::Excluded(val) => ( - calc_fraction(&bucket.lower, &bucket.upper, val)?, + encoded_fraction(&bucket.lower, &bucket.upper, val, true, buf)?, endpoint_count(val, bucket, sketch), ), Bound::Unbounded => unreachable!(), @@ -577,11 +530,12 @@ impl Histogram { *bucket_i += 1; } else if is_under(comparator, &bucket.upper, max, false)? { let (temp_ratio, option) = match min { - Bound::Included(val) => { - (calc_fraction(&bucket.lower, &bucket.upper, val)?, None) - } + Bound::Included(val) => ( + encoded_fraction(&bucket.lower, &bucket.upper, val, false, buf)?, + None, + ), Bound::Excluded(val) => ( - calc_fraction(&bucket.lower, &bucket.upper, val)?, + encoded_fraction(&bucket.lower, &bucket.upper, val, false, buf)?, endpoint_count(val, bucket, sketch), ), Bound::Unbounded => unreachable!(), @@ -594,21 +548,23 @@ impl Histogram { *bucket_i += 1; } else { let (temp_ratio_max, option_max) = match max { - Bound::Included(val) => { - (calc_fraction(&bucket.lower, &bucket.upper, val)?, None) - } + Bound::Included(val) => ( + encoded_fraction(&bucket.lower, &bucket.upper, val, true, buf)?, + None, + ), Bound::Excluded(val) => ( - calc_fraction(&bucket.lower, &bucket.upper, val)?, + encoded_fraction(&bucket.lower, &bucket.upper, val, true, buf)?, endpoint_count(val, bucket, sketch), ), Bound::Unbounded => unreachable!(), }; let (temp_ratio_min, option_min) = match min { - Bound::Included(val) => { - (calc_fraction(&bucket.lower, &bucket.upper, val)?, None) - } + Bound::Included(val) => ( + encoded_fraction(&bucket.lower, &bucket.upper, val, false, buf)?, + None, + ), Bound::Excluded(val) => ( - calc_fraction(&bucket.lower, &bucket.upper, val)?, + encoded_fraction(&bucket.lower, &bucket.upper, val, false, buf)?, endpoint_count(val, bucket, sketch), ), Bound::Unbounded => unreachable!(), @@ -642,6 +598,44 @@ impl Histogram { } } +fn encoded_fraction( + start: &DataValue, + end: &DataValue, + value: &DataValue, + is_upper: bool, + buf: &mut Vec, +) -> Result { + buf.clear(); + start.memcomparable_encode(buf)?; + let lower_end = buf.len(); + end.memcomparable_encode(buf)?; + let upper_end = buf.len(); + value.memcomparable_encode(buf)?; + if is_upper && matches!(value, DataValue::Tuple(_)) { + buf.push(crate::storage::table_codec::BOUND_MAX_TAG); + } + let lower = &buf[..lower_end]; + let upper = &buf[lower_end..upper_end]; + let key = &buf[upper_end..]; + if key <= lower { + return Ok(0.0); + } + if key >= upper { + return Ok(1.0); + } + let prefix = lower.iter().zip(upper).take_while(|(a, b)| a == b).count(); + // Subtract eight-byte integer coordinates before converting to f64, so + // nearby large keys retain their distance. Encoded distance is approximate. + let coordinate = |bytes: &[u8]| { + (0..8).fold(0u64, |value, i| { + (value << 8) | u64::from(bytes.get(prefix + i).copied().unwrap_or(0)) + }) + }; + let lower = coordinate(lower); + let upper = coordinate(upper); + Ok((coordinate(key) - lower) as f64 / (upper - lower) as f64) +} + fn subtract_endpoint_count(count: usize, endpoint_count: usize) -> usize { if endpoint_count < count { count.saturating_sub(endpoint_count) @@ -661,12 +655,7 @@ fn endpoint_count( debug_assert_eq!(bucket_key_type, bucket.upper.logical_type()); if value.logical_type() == bucket_key_type { - match value { - DataValue::Tuple(values, true) => { - Some(sketch.estimate(&DataValue::Tuple(values.clone(), false))) - } - _ => Some(sketch.estimate(value)), - } + Some(sketch.estimate(value)) } else { None } @@ -913,6 +902,314 @@ mod tests { Ok(()) } + #[test] + fn parameterized_tuple_ranges_accumulate_prefix_selectivity() -> Result<(), DatabaseError> { + let mut builder = HistogramBuilder::new(&index_meta(), ANALYZE_STATISTICS_RELATIVE_ERROR)?; + for value in 0..100 { + builder.append(DataValue::Tuple(vec![DataValue::Int32(value); 4]))?; + } + let (mut histogram, sketch, top_n) = builder.build(10)?; + // Fixed metadata isolates the heuristic from HLL estimation error. + histogram.meta.values_len = 10_000; + histogram.meta.number_of_distinct_value = 10_000; + let parameter = |id| DataValue::Parameter { + id, + ty: crate::types::LogicalType::Integer, + }; + let prefix = vec![parameter(1), parameter(2)]; + let scope = |lower, upper| Range::Scope { + min: Bound::Excluded(DataValue::Tuple(lower)), + max: Bound::Excluded(DataValue::Tuple(upper)), + }; + assert_eq!( + histogram.collect_count(&[scope(prefix.clone(), prefix.clone())], &sketch, &top_n)?, + 100 + ); + let longer = vec![parameter(1), parameter(2), parameter(3)]; + assert_eq!( + histogram.collect_count(&[scope(longer.clone(), longer)], &sketch, &top_n)?, + 10 + ); + let mut lower = prefix.clone(); + lower.push(parameter(3)); + let mut upper = prefix.clone(); + upper.push(parameter(4)); + assert_eq!( + histogram.collect_count(&[scope(lower.clone(), prefix)], &sketch, &top_n)?, + 25 + ); + assert_eq!( + histogram.collect_count(&[scope(lower.clone(), upper.clone())], &sketch, &top_n)?, + 2 + ); + let full = DataValue::Tuple(vec![parameter(1), parameter(2), parameter(3), parameter(4)]); + assert_eq!( + histogram.collect_count(&[Range::Eq(full.clone())], &sketch, &top_n)?, + 1 + ); + assert_eq!( + histogram.collect_count( + &[Range::Scope { + min: Bound::Excluded(full.clone()), + max: Bound::Included(full), + }], + &sketch, + &top_n + )?, + 0 + ); + // Equal trailing elements after the first differing element are not a prefix. + lower.push(parameter(5)); + upper.push(parameter(5)); + assert_eq!( + histogram.collect_count(&[scope(lower, upper)], &sketch, &top_n)?, + 2 + ); + histogram.buckets[0].upper = DataValue::Tuple(vec![ + DataValue::Tuple(vec![DataValue::Int32(0); 2]), + DataValue::Tuple(vec![DataValue::Int32(0); 2]), + ]); + let nested = vec![DataValue::Tuple(vec![parameter(1), parameter(2)])]; + assert_eq!( + histogram.collect_count(&[scope(nested.clone(), nested)], &sketch, &top_n)?, + 100 + ); + Ok(()) + } + + #[test] + fn encoded_fraction_expected_ratios_by_type() -> Result<(), DatabaseError> { + use crate::types::value::Utf8Type; + use crate::types::CharLengthUnits; + + #[allow(unused_mut)] + let mut cases: Vec<(&str, Vec)> = vec![ + ("null", vec![DataValue::Null]), + ( + "boolean", + vec![DataValue::Boolean(false), DataValue::Boolean(true)], + ), + ("i8", [-100, -50, 0, 50, 100].map(DataValue::Int8).to_vec()), + ( + "i16", + [-100, -50, 0, 50, 100].map(DataValue::Int16).to_vec(), + ), + ( + "i32", + [-100, -50, 0, 50, 100].map(DataValue::Int32).to_vec(), + ), + ( + "i64", + [-100, -50, 0, 50, 100].map(DataValue::Int64).to_vec(), + ), + ("u8", [0, 50, 100, 150, 200].map(DataValue::UInt8).to_vec()), + ( + "u16", + [0, 50, 100, 150, 200].map(DataValue::UInt16).to_vec(), + ), + ( + "u32", + [0, 50, 100, 150, 200].map(DataValue::UInt32).to_vec(), + ), + ( + "u64", + [0, 50, 100, 150, 200].map(DataValue::UInt64).to_vec(), + ), + ( + "f32", + [1.0, 1.25, 1.5, 1.75, 2.0] + .map(|v| DataValue::Float32(v.into())) + .to_vec(), + ), + ( + "f64", + [1.0, 1.25, 1.5, 1.75, 2.0] + .map(|v| DataValue::Float64(v.into())) + .to_vec(), + ), + ( + "date", + [0, 50, 100, 150, 200].map(DataValue::Date32).to_vec(), + ), + ( + "datetime", + [0, 50, 100, 150, 200].map(DataValue::Date64).to_vec(), + ), + ( + "time32", + [0, 50, 100, 150, 200] + .map(|v| DataValue::Time32(v, 0)) + .to_vec(), + ), + ( + "time64", + [0, 50, 100, 150, 200] + .map(|v| DataValue::Time64(v, 6, false)) + .to_vec(), + ), + ( + "timestamp", + [0, 50, 100, 150, 200] + .map(|v| DataValue::Time64(v, 6, true)) + .to_vec(), + ), + ( + "varchar", + ["a", "b", "c", "d", "e"] + .map(|v| DataValue::from(v.to_string())) + .to_vec(), + ), + ( + "char", + ["a", "b", "c", "d", "e"] + .map(|v| DataValue::Utf8 { + value: v.into(), + ty: Utf8Type::Fixed(1), + unit: CharLengthUnits::Characters, + }) + .to_vec(), + ), + ( + "unicode", + ["一", "丁", "丂", "七", "丄"] + .map(|v| DataValue::from(v.to_string())) + .to_vec(), + ), + ( + "long string", + ["a", "b", "c", "d", "e"] + .map(|v| DataValue::from(format!("{}{}", "prefix".repeat(20), v))) + .to_vec(), + ), + ( + "tuple", + [0, 50, 100, 150, 200] + .map(|v| DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Int32(v)])) + .to_vec(), + ), + ( + "nested tuple", + [0, 50, 100, 150, 200] + .map(|v| DataValue::Tuple(vec![DataValue::Tuple(vec![DataValue::Int32(v)])])) + .to_vec(), + ), + ( + "i64 max", + (0..5).map(|v| DataValue::Int64(i64::MAX - 4 + v)).collect(), + ), + ( + "u64 max", + (0..5) + .map(|v| DataValue::UInt64(u64::MAX - 4 + v)) + .collect(), + ), + ]; + #[cfg(feature = "decimal")] + cases.push(( + "decimal", + [100, 125, 150, 175, 200] + .map(|v| DataValue::Decimal(rust_decimal::Decimal::new(v, 2))) + .to_vec(), + )); + let mut buf = Vec::new(); + for (name, values) in &cases { + let expected: &[f64] = match *name { + "null" => &[0.0], + "boolean" => &[0.0, 1.0], + // Accepted encoded-space ratios, not numeric-space quartiles. + "decimal" => &[0.0, 0.59765625, 0.6953125, 0.79296875, 1.0], + _ => &[0.0, 0.25, 0.5, 0.75, 1.0], + }; + let lower = values.first().unwrap(); + let upper = values.last().unwrap(); + let fractions = values + .iter() + .map(|value| super::encoded_fraction(lower, upper, value, false, &mut buf)) + .collect::, _>>()?; + assert_eq!(fractions.len(), expected.len(), "{name}"); + for (i, (actual, expected)) in fractions.iter().zip(expected).enumerate() { + assert!( + (actual - expected).abs() <= 1e-12, + "{name}[{i}]: expected {expected}, got {actual}" + ); + } + } + Ok(()) + } + + #[test] + fn tuple_interpolation_uses_encoded_order_and_prefix_markers() -> Result<(), DatabaseError> { + let tuple = |a, b| DataValue::Tuple(vec![DataValue::Int32(a), DataValue::Int32(b)]); + let mut buf = Vec::new(); + let fraction = + super::encoded_fraction(&tuple(1, 0), &tuple(2, 0), &tuple(1, 100), false, &mut buf)?; + assert!((fraction - 100.0 / 256.0_f64.powi(5)).abs() <= 1e-20); + let lower = DataValue::Tuple(vec![DataValue::Int32(1)]); + let upper = DataValue::Tuple(vec![DataValue::Int32(1)]); + assert_eq!( + super::encoded_fraction(&tuple(1, 0), &tuple(2, 0), &lower, false, &mut buf)?, + 0.0 + ); + assert_eq!( + super::encoded_fraction(&tuple(1, 0), &tuple(2, 0), &upper, true, &mut buf)?, + 0.994140625 + ); + let string_tuple = + |s: &str| DataValue::Tuple(vec![DataValue::Int32(1), DataValue::from(s.to_string())]); + let fraction = super::encoded_fraction( + &string_tuple("Alice"), + &string_tuple("Zoe"), + &string_tuple("Bob"), + false, + &mut buf, + )?; + assert!((fraction - 0.04044539011559242).abs() <= 1e-12); + Ok(()) + } + + #[test] + fn parameterized_range_count_uses_shape_and_ndv() -> Result<(), DatabaseError> { + let mut builder = HistogramBuilder::new(&index_meta(), ANALYZE_STATISTICS_RELATIVE_ERROR)?; + for value in 0..1_000 { + builder.append(DataValue::Int32(value % 100))?; + } + let (histogram, sketch, top_n) = builder.build(10)?; + let parameter = |id| DataValue::Parameter { + id, + ty: crate::types::LogicalType::Integer, + }; + for (range, expected) in [ + ( + Range::Eq(parameter(1)), + 1_000usize.div_ceil(histogram.distinct_values_len()), + ), + ( + Range::Scope { + min: Bound::Included(parameter(1)), + max: Bound::Unbounded, + }, + 250, + ), + ( + Range::Scope { + min: Bound::Included(parameter(1)), + max: Bound::Excluded(parameter(2)), + }, + 16, + ), + ( + Range::SortedRanges(vec![Range::Eq(parameter(1)), Range::Eq(parameter(2))]), + 2 * 1_000usize.div_ceil(histogram.distinct_values_len()), + ), + ] { + assert_eq!( + histogram.collect_count(&[range], &sketch, &top_n)?, + expected + ); + } + Ok(()) + } + #[test] fn test_collect_count() -> Result<(), DatabaseError> { let mut builder = HistogramBuilder::new(&index_meta(), ANALYZE_STATISTICS_RELATIVE_ERROR)?; @@ -1216,21 +1513,21 @@ mod tests { let mut builder = HistogramBuilder::new(&index_meta(), ANALYZE_STATISTICS_RELATIVE_ERROR)?; for value in 0..15 { - builder.append(DataValue::Tuple( - vec![DataValue::Int32(value), DataValue::Int32(value)], - false, - ))?; + builder.append(DataValue::Tuple(vec![ + DataValue::Int32(value), + DataValue::Int32(value), + ]))?; } let (histogram, mut sketch, top_n) = builder.build(5)?; let ranges = [Range::Scope { - min: Bound::Excluded(DataValue::Tuple(vec![DataValue::Int32(0)], false)), - max: Bound::Excluded(DataValue::Tuple(vec![DataValue::Int32(8)], true)), + min: Bound::Excluded(DataValue::Tuple(vec![DataValue::Int32(0)])), + max: Bound::Excluded(DataValue::Tuple(vec![DataValue::Int32(8)])), }]; let clean_count = histogram.collect_count(&ranges, &sketch, &top_n)?; - sketch.increment(&DataValue::Tuple(vec![DataValue::Int32(0)], false)); - sketch.increment(&DataValue::Tuple(vec![DataValue::Int32(8)], true)); + sketch.increment(&DataValue::Tuple(vec![DataValue::Int32(0)])); + sketch.increment(&DataValue::Tuple(vec![DataValue::Int32(8)])); assert_eq!( histogram.collect_count(&ranges, &sketch, &top_n)?, @@ -1243,8 +1540,8 @@ mod tests { #[test] fn test_endpoint_count_uses_only_full_histogram_keys() -> Result<(), DatabaseError> { let bucket = Bucket { - lower: DataValue::Tuple(vec![DataValue::Int32(0), DataValue::Int32(0)], false), - upper: DataValue::Tuple(vec![DataValue::Int32(10), DataValue::Int32(10)], false), + lower: DataValue::Tuple(vec![DataValue::Int32(0), DataValue::Int32(0)]), + upper: DataValue::Tuple(vec![DataValue::Int32(10), DataValue::Int32(10)]), count: 11, }; let mut sketch = CountMinSketch::with_relative_error( @@ -1252,9 +1549,9 @@ mod tests { ANALYZE_STATISTICS_RELATIVE_ERROR, )?; - let real_key = DataValue::Tuple(vec![DataValue::Int32(8), DataValue::Int32(8)], false); - let upper_bound = DataValue::Tuple(vec![DataValue::Int32(8), DataValue::Int32(8)], true); - let prefix_bound = DataValue::Tuple(vec![DataValue::Int32(8)], true); + let real_key = DataValue::Tuple(vec![DataValue::Int32(8), DataValue::Int32(8)]); + let upper_bound = DataValue::Tuple(vec![DataValue::Int32(8), DataValue::Int32(8)]); + let prefix_bound = DataValue::Tuple(vec![DataValue::Int32(8)]); sketch.increment(&real_key); sketch.increment(&prefix_bound); diff --git a/src/optimizer/rule/normalization/column_pruning.rs b/src/optimizer/rule/normalization/column_pruning.rs index 4eae5a7d..a17a1007 100644 --- a/src/optimizer/rule/normalization/column_pruning.rs +++ b/src/optimizer/rule/normalization/column_pruning.rs @@ -25,6 +25,7 @@ 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::MetaArena; use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; @@ -90,7 +91,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, @@ -103,7 +104,7 @@ impl ExprVisitor> for ReferencedColumnCollector<'_, '_> { &mut self, expr: ExprRef, _ty: &AliasType, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), ) -> Result<(), DatabaseError> { self.visit(expr, arena) } @@ -129,7 +130,8 @@ impl ColumnPruning { referenced_columns, arena, }; - OperatorExprVisitor::new(&mut collector, arena).visit_operator(operator)?; + OperatorExprVisitor::new(&mut collector, arena as &dyn MetaArena) + .visit_operator(operator)?; struct ReferencedOperatorColumnCollector<'a, 'p> { referenced_columns: &'a mut ReferencedColumns, @@ -268,7 +270,7 @@ impl ColumnPruning { &mut PositionRemapper::new(removed_positions, remapped_exprs), arena, ) - .visit_operator(operator) + .visit_operator(operator, None) } fn remap_exprs_after_child_change<'a>( diff --git a/src/optimizer/rule/normalization/combine_operators.rs b/src/optimizer/rule/normalization/combine_operators.rs index dd28d35d..f432d891 100644 --- a/src/optimizer/rule/normalization/combine_operators.rs +++ b/src/optimizer/rule/normalization/combine_operators.rs @@ -13,7 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; -use crate::expression::visitor_mut::ExprVisitorMut; +use crate::expression::visitor_mut::{walk_mut_expr, ExprVisitorMut}; use crate::expression::{BinaryOperator, ScalarExpression}; use crate::optimizer::core::rule::NormalizationRule; use crate::optimizer::plan_utils::{only_child_mut, replace_with_only_child}; @@ -21,6 +21,7 @@ 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::MetaArena; use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::LogicalType; use std::mem; @@ -59,7 +60,7 @@ fn rewrite_column_position( &mut self, _column: &mut crate::catalog::ColumnRef, position: &mut usize, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { *position = self.0; Ok(()) @@ -149,6 +150,75 @@ impl NormalizationRule for CollapseProject { } } +// Removing a project must preserve every input slot; renaming is allowed. +fn can_remove_project( + op: &ProjectOperator, + childrens: &mut Childrens, + arena: &mut PlanArena<'_>, +) -> bool { + let Childrens::Only(input) = childrens else { + return false; + }; + let schema = input.output_schema(arena); + if op.exprs.len() != schema.len() { + return false; + } + for (i, expr) in op.exprs.iter().enumerate() { + let mut source_expr = *expr; + while let ScalarExpression::Alias { expr, .. } = arena.expression(source_expr) { + source_expr = *expr; + } + let ScalarExpression::ColumnRef { + column: source_column, + .. + } = arena.expression(source_expr) + else { + return false; + }; + let ScalarExpression::ColumnRef { + position: execution_position, + .. + } = arena.expression(strip_alias(*expr, arena)) + else { + return false; + }; + if *source_column != schema[i] || *execution_position != i { + return false; + } + } + true +} + +struct ProjectionSubstitution<'a>(&'a [ExprRef]); + +impl ExprVisitorMut for ProjectionSubstitution<'_> { + fn visit( + &mut self, + expr: &mut ExprRef, + arena: &mut (dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { + if let ScalarExpression::ColumnRef { position, .. } = arena.expression(*expr) { + let Some(source) = self.0.get(*position) else { + return Err(DatabaseError::InvalidValue( + "invalid projection slot".into(), + )); + }; + *expr = source.clone_expression(arena)?; + return Ok(()); + } + walk_mut_expr(self, expr, arena) + } + + fn visit_alias( + &mut self, + expr: &mut ExprRef, + _: &mut crate::expression::AliasType, + arena: &mut (dyn MetaArena + '_), + ) -> Result<(), DatabaseError> { + self.visit(expr, arena) + } +} + /// Combine two adjacent filter operators into one. pub struct CombineFilter; @@ -165,7 +235,7 @@ impl NormalizationRule for CombineFilter { return Ok(false); } }; - let parent_filter = parent_filter; + let mut parent_filter = parent_filter; let cursor = match only_child_mut(plan) { Some(child) => child, @@ -175,6 +245,7 @@ impl NormalizationRule for CombineFilter { } }; + let mut changed = false; loop { match &mut cursor.operator { Operator::Filter(child_op) => { @@ -192,11 +263,21 @@ impl NormalizationRule for CombineFilter { ty: LogicalType::Boolean, }); child_op.having = having || child_op.having; + child_op.is_optimized = false; return Ok(replace_with_only_child(plan)); } - Operator::Project(project_op) if is_passthrough_project(project_op, arena) => { + Operator::Project(project_op) => { + if !can_remove_project(project_op, &mut cursor.childrens, arena) { + plan.operator = Operator::Filter(parent_filter); + return Ok(changed); + } + let mut predicate = parent_filter.predicate.clone_expression(arena)?; + ProjectionSubstitution(&project_op.exprs).visit(&mut predicate, arena)?; + parent_filter.predicate = predicate; + parent_filter.is_optimized = false; if replace_with_only_child(cursor) { + changed = true; continue; } plan.operator = Operator::Filter(parent_filter); @@ -204,7 +285,7 @@ impl NormalizationRule for CombineFilter { } _ => { plan.operator = Operator::Filter(parent_filter); - return Ok(false); + return Ok(changed); } } } @@ -379,6 +460,25 @@ mod tests { Ok(()) } + #[test] + fn test_combine_filter_keeps_reordered_projection() -> Result<(), DatabaseError> { + let table_state = build_t1_table()?; + let mut arena = PlanArena::new(&table_state.table_arena); + let plan = table_state.plan_with_arena("select * from t1 c where c.c1 > 1", &mut arena)?; + let mut filter = plan.childrens.pop_only(); + if let Childrens::Only(child) = filter.childrens.as_mut() { + if let Operator::Project(project) = &mut child.operator { + project.exprs.reverse(); + } + } + assert!(!super::CombineFilter.apply(&mut filter, &mut arena)?); + assert!(matches!( + filter.childrens.pop_only().operator, + Operator::Project(_) + )); + Ok(()) + } + #[test] fn test_combine_filter() -> Result<(), DatabaseError> { let table_state = build_t1_table()?; diff --git a/src/optimizer/rule/normalization/compilation_in_advance.rs b/src/optimizer/rule/normalization/compilation_in_advance.rs index 08a909cc..e35a8c19 100644 --- a/src/optimizer/rule/normalization/compilation_in_advance.rs +++ b/src/optimizer/rule/normalization/compilation_in_advance.rs @@ -27,7 +27,7 @@ pub(crate) fn evaluator_bind_current( arena: &mut PlanArena, ) -> Result<(), DatabaseError> { let mut evaluator = BindEvaluator; - OperatorExprVisitorMut::new(&mut evaluator, arena).visit_operator(&mut plan.operator) + OperatorExprVisitorMut::new(&mut evaluator, arena).visit_operator(&mut plan.operator, None) } impl EvaluatorBind { diff --git a/src/optimizer/rule/normalization/elimination.rs b/src/optimizer/rule/normalization/elimination.rs index 8ff77852..045a0e45 100644 --- a/src/optimizer/rule/normalization/elimination.rs +++ b/src/optimizer/rule/normalization/elimination.rs @@ -770,11 +770,35 @@ 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, c2.clone()], 1); - super::mark_sort_preserving_indexes(&mut plan, &[c2], &arena)?; - let rule = EliminateRedundantSort; - - assert!(rule.apply(&mut plan, &mut arena)?); + for lookup in [ + None, + Some(IndexLookup::Static(Range::Eq(DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + }))), + ] { + let mut plan = build_plan( + &mut arena, + vec![c2.clone()], + vec![c1.clone(), c2.clone()], + 1, + ); + let Childrens::Only(filter) = plan.childrens.as_mut() else { + panic!("expected filter") + }; + let Childrens::Only(leaf) = filter.childrens.as_mut() else { + panic!("expected index scan") + }; + if let Some(PhysicalOption { + plan: PlanImpl::IndexScan(info), + .. + }) = &mut leaf.physical_option + { + info.lookup = lookup; + } + super::mark_sort_preserving_indexes(&mut plan, &[c2.clone()], &arena)?; + assert!(EliminateRedundantSort.apply(&mut plan, &mut arena)?); + } Ok(()) } diff --git a/src/optimizer/rule/normalization/mod.rs b/src/optimizer/rule/normalization/mod.rs index b85d4118..f6ff5ad8 100644 --- a/src/optimizer/rule/normalization/mod.rs +++ b/src/optimizer/rule/normalization/mod.rs @@ -22,6 +22,7 @@ use crate::optimizer::rule::normalization::combine_operators::{ }; use crate::optimizer::rule::normalization::compilation_in_advance::EvaluatorBind; use crate::planner::operator::Operator; +use crate::planner::MetaArena; use crate::optimizer::rule::normalization::min_max_top_k::MinMaxToTopK; use crate::optimizer::rule::normalization::pushdown_limit::{ @@ -280,7 +281,7 @@ impl ExprVisitorMut for PositionRemapper<'_, '_> { fn visit_expression_ref( &mut self, expr: &mut ExprRef, - _arena: &mut crate::planner::PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result { Ok(self.visited.insert(*expr)) } @@ -289,7 +290,7 @@ impl ExprVisitorMut for PositionRemapper<'_, '_> { &mut self, _column: &mut crate::catalog::ColumnRef, position: &mut usize, - _arena: &mut crate::planner::PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { remap_position(position, self.removed_positions); Ok(()) @@ -299,7 +300,7 @@ impl ExprVisitorMut for PositionRemapper<'_, '_> { &mut self, expr: &mut ExprRef, alias: &mut AliasType, - arena: &mut crate::planner::PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { match alias { AliasType::Expr(alias_expr) => self.visit(alias_expr, arena), diff --git a/src/optimizer/rule/normalization/parameterized_index.rs b/src/optimizer/rule/normalization/parameterized_index.rs index 4ea2eacb..bc726d29 100644 --- a/src/optimizer/rule/normalization/parameterized_index.rs +++ b/src/optimizer/rule/normalization/parameterized_index.rs @@ -23,6 +23,7 @@ use crate::planner::operator::mark_apply::{MarkApplyKind, MarkApplyOperator, Mar use crate::planner::operator::project::ProjectOperator; use crate::planner::operator::table_scan::TableScanOperator; use crate::planner::operator::{Operator, PhysicalOption, PlanImpl, SortOption}; +use crate::planner::MetaArena; use crate::planner::{Childrens, ExprRef, LogicalPlan, PlanArena}; use crate::types::index::{IndexLookup, IndexType}; use crate::types::tuple::Schema; @@ -66,7 +67,7 @@ fn find_parameterized_probe( predicates: &[ExprRef], left_schema: &Schema, right_schema: &Schema, - arena: &crate::planner::PlanArena, + arena: &dyn MetaArena, ) -> Result, DatabaseError> { match kind { MarkApplyKind::Exists => { @@ -94,7 +95,7 @@ fn extract_parameterized_probe( predicate: ExprRef, left_schema: &Schema, right_schema: &Schema, - arena: &crate::planner::PlanArena, + arena: &dyn MetaArena, ) -> Result, DatabaseError> { match predicate.unpack_alias_ref(arena) { ScalarExpression::Binary { @@ -129,7 +130,7 @@ fn extract_parameterized_probe_side( left_expr: ExprRef, left_schema: &Schema, right_schema: &Schema, - arena: &crate::planner::PlanArena, + arena: &dyn MetaArena, ) -> Result, DatabaseError> { let Some((right_column, _)) = right_expr .unpack_alias(arena) @@ -158,7 +159,7 @@ fn extract_parameterized_probe_side( fn parameterize_right_subtree( plan: &mut LogicalPlan, right_column: &ColumnRef, - arena: &crate::planner::PlanArena, + arena: &dyn MetaArena, ) -> bool { if matches!(plan.operator, Operator::TableScan(_)) { let index_info = { @@ -203,7 +204,7 @@ fn parameterize_right_subtree( fn pick_parameterized_index_position( scan_op: &TableScanOperator, right_column: &ColumnRef, - arena: &crate::planner::PlanArena, + arena: &dyn MetaArena, ) -> Option { let right_column = arena.column(*right_column); let column_id = right_column.id()?; @@ -235,11 +236,7 @@ fn index_priority(index_type: IndexType) -> usize { } } -fn schema_contains_column( - schema: &Schema, - column: &ColumnRef, - arena: &crate::planner::PlanArena, -) -> bool { +fn schema_contains_column(schema: &Schema, column: &ColumnRef, arena: &dyn MetaArena) -> bool { schema .iter() .any(|candidate| arena.same_column(*candidate, *column)) diff --git a/src/optimizer/rule/normalization/pushdown_predicates.rs b/src/optimizer/rule/normalization/pushdown_predicates.rs index 637c8267..8eb9ade8 100644 --- a/src/optimizer/rule/normalization/pushdown_predicates.rs +++ b/src/optimizer/rule/normalization/pushdown_predicates.rs @@ -406,9 +406,9 @@ impl PushPredicateIntoScan { if range.only_eq() && apply_column_count != column_count { fn eq_to_scope(range: Range) -> Range { match range { - Range::Eq(DataValue::Tuple(values, _)) => { - let min = Bound::Included(DataValue::Tuple(values.clone(), false)); - let max = Bound::Included(DataValue::Tuple(values, true)); + Range::Eq(DataValue::Tuple(values)) => { + let min = Bound::Included(DataValue::Tuple(values.clone())); + let max = Bound::Included(DataValue::Tuple(values)); Range::Scope { min, max } } @@ -695,14 +695,11 @@ mod tests { assert_eq!(ignore_prefix_len, 3); assert_eq!( detached.range, - Range::Eq(DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(3), - ], - false, - )) + Range::Eq(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2), + DataValue::Int32(3), + ],)) ); let residual = detached.residual.expect("c4 predicate should remain"); let residual_detached = RangeDetacher::new(table_name.as_ref(), &4, &mut arena) @@ -760,11 +757,11 @@ mod tests { assert_eq!( detached.range, Range::Scope { - min: Bound::Excluded(DataValue::Tuple( - vec![DataValue::Int32(1), DataValue::Int32(2)], - true, - )), - max: Bound::Excluded(DataValue::Tuple(vec![DataValue::Int32(1)], true)), + min: Bound::Excluded(DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Int32(2) + ],)), + max: Bound::Included(DataValue::Tuple(vec![DataValue::Int32(1)])), } ); let residual = detached.residual.expect("c3 predicate should remain"); diff --git a/src/optimizer/rule/normalization/simplification.rs b/src/optimizer/rule/normalization/simplification.rs index 604d37ac..6b9276b2 100644 --- a/src/optimizer/rule/normalization/simplification.rs +++ b/src/optimizer/rule/normalization/simplification.rs @@ -28,7 +28,7 @@ pub(crate) fn constant_calculation_current( arena: &mut crate::planner::PlanArena, ) -> Result<(), DatabaseError> { let mut calculator = ConstantCalculator::new(arena); - OperatorExprVisitorMut::new(&mut calculator, arena).visit_operator(&mut plan.operator) + OperatorExprVisitorMut::new(&mut calculator, arena).visit_operator(&mut plan.operator, None) } impl ConstantCalculation { @@ -419,6 +419,50 @@ mod test { Ok(()) } + #[test] + fn test_non_comparison_does_not_apply_pending_rearrangement() -> Result<(), DatabaseError> { + let table_state = build_t1_table()?; + let mut arena = PlanArena::new(&table_state.table_arena); + let plan = + table_state.plan_with_arena("select * from t1 where (c1 + 1) % 2 = 0", &mut arena)?; + + let best_plan = run_with_single_batch( + plan, + "test_non_comparison_does_not_apply_pending_rearrangement", + HepBatchStrategy::once_topdown(), + vec![NormalizationRuleImpl::SimplifyFilter], + &mut arena, + )?; + let filter = best_plan.childrens.pop_only(); + let Operator::Filter(filter) = filter.operator else { + panic!("expected filter"); + }; + let ScalarExpression::Binary { + op: BinaryOperator::Eq, + left_expr, + .. + } = arena.expression(filter.predicate) + else { + panic!("expected equality"); + }; + let ScalarExpression::Binary { + op: BinaryOperator::Modulo, + left_expr, + .. + } = arena.expression(*left_expr) + else { + panic!("expected modulo"); + }; + assert!(matches!( + arena.expression(*left_expr), + ScalarExpression::Binary { + op: BinaryOperator::Plus, + .. + } + )); + Ok(()) + } + fn plan_filter( plan: &LogicalPlan, column_id: &ColumnId, diff --git a/src/orm/ddl.rs b/src/orm/ddl.rs index 322227bb..1c48eea5 100644 --- a/src/orm/ddl.rs +++ b/src/orm/ddl.rs @@ -96,7 +96,7 @@ impl Database { 'parent, 'arena, S::TransactionType<'_>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { @@ -117,7 +117,7 @@ impl Database { 'parent, 'arena, S::TransactionType<'_>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { @@ -487,11 +487,11 @@ where 'parent, 'arena, S::TransactionType<'_>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { - static EMPTY_ORM_PARAMS: &[(&str, DataValue)] = &[]; + static EMPTY_ORM_PARAMS: &[(usize, LogicalType)] = &[]; let view_name = view_name.to_string(); database.execute_mut("ORM CREATE VIEW", EMPTY_ORM_PARAMS, move |binder, arena| { let mut context = OrmContext { binder, arena }; diff --git a/src/orm/dql.rs b/src/orm/dql.rs index 0a8e1865..d6292246 100644 --- a/src/orm/dql.rs +++ b/src/orm/dql.rs @@ -11,7 +11,7 @@ impl Database { 'parent, 'arena, S::TransactionType<'_>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { @@ -28,7 +28,7 @@ impl Database { 'parent, 'arena, S::TransactionType<'_>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { @@ -170,7 +170,7 @@ impl<'a, S: Storage> DBTransaction<'a, S> { 'parent, 'arena, S::TransactionType<'a>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { @@ -187,7 +187,7 @@ impl<'a, S: Storage> DBTransaction<'a, S> { 'parent, 'arena, S::TransactionType<'a>, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { diff --git a/src/orm/mod.rs b/src/orm/mod.rs index 8b1aae23..ce49171c 100644 --- a/src/orm/mod.rs +++ b/src/orm/mod.rs @@ -220,7 +220,7 @@ impl FieldSort { pub trait BindOrmScalar<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_scalar( self, @@ -231,7 +231,7 @@ where impl<'bind, 'parent, 'arena, T, A, M, V> BindOrmScalar<'bind, 'parent, 'arena, T, A> for Field where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_scalar( self, @@ -244,7 +244,7 @@ where impl<'bind, 'parent, 'arena, T, A> BindOrmScalar<'bind, 'parent, 'arena, T, A> for ScalarExpression where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_scalar( self, @@ -258,7 +258,7 @@ impl<'bind, 'parent, 'arena, T, A> BindOrmScalar<'bind, 'parent, 'arena, T, A> for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_scalar( self, @@ -272,7 +272,7 @@ where pub trait BindOrmSort<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -283,7 +283,7 @@ where impl<'bind, 'parent, 'arena, T, A, M, V> BindOrmSort<'bind, 'parent, 'arena, T, A> for Field where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -297,7 +297,7 @@ impl<'bind, 'parent, 'arena, T, A, M, V> BindOrmSort<'bind, 'parent, 'arena, T, for FieldSort where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -313,7 +313,7 @@ where impl<'bind, 'parent, 'arena, T, A> BindOrmSort<'bind, 'parent, 'arena, T, A> for ScalarExpression where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -326,7 +326,7 @@ where impl<'bind, 'parent, 'arena, T, A> BindOrmSort<'bind, 'parent, 'arena, T, A> for SortField where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -340,7 +340,7 @@ impl<'bind, 'parent, 'arena, T, A> BindOrmSort<'bind, 'parent, 'arena, T, A> for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_sort<'scope>( self, @@ -380,7 +380,7 @@ impl WindowSpec { impl<'bind, 'parent, 'arena, T, A> From> for ExprRef where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn from(expr: CtxExpression<'bind, 'parent, 'arena, T, A>) -> Self { expr.into_scalar() @@ -418,7 +418,7 @@ where pub trait BindOrmScalarList<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn bind_scalar_list( self, @@ -433,7 +433,7 @@ macro_rules! impl_bind_orm_scalar_list { for ($($name,)+) where Tx: Transaction, - Args: AsRef<[(&'static str, DataValue)]>, + Args: AsRef<[(usize, LogicalType)]>, $($name: BindOrmScalar<'bind, 'parent, 'arena, Tx, Args>,)+ { #[allow(non_snake_case)] @@ -485,7 +485,7 @@ macro_rules! impl_quantified_subquery_methods { struct ExprBindScopeHandle<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: NonNull>, arena: NonNull>, @@ -495,7 +495,7 @@ where impl<'bind, 'parent, 'arena, T, A> Clone for ExprBindScopeHandle<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn clone(&self) -> Self { *self @@ -505,14 +505,14 @@ where impl<'bind, 'parent, 'arena, T, A> Copy for ExprBindScopeHandle<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { } impl<'bind, 'parent, 'arena, T, A> ExprBindScopeHandle<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn new<'ctx>(scope: &ExprBindScope<'ctx, 'bind, 'parent, 'arena, T, A>) -> Self { Self { @@ -689,7 +689,7 @@ where pub struct CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { expr: ExprRef, scope: ExprBindScopeHandle<'bind, 'parent, 'arena, T, A>, @@ -698,7 +698,7 @@ where impl<'bind, 'parent, 'arena, T, A> CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub fn into_scalar(self) -> ExprRef { self.expr @@ -963,7 +963,7 @@ where impl<'bind, 'parent, 'arena, T, A> fmt::Debug for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { self.expr.fmt(f) @@ -973,7 +973,7 @@ where impl<'bind, 'parent, 'arena, T, A> Clone for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn clone(&self) -> Self { Self { @@ -986,7 +986,7 @@ where impl<'bind, 'parent, 'arena, T, A> PartialEq for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn eq(&self, other: &Self) -> bool { self.expr == other.expr @@ -996,14 +996,14 @@ where impl<'bind, 'parent, 'arena, T, A> Eq for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { } impl<'bind, 'parent, 'arena, T, A> Hash for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn hash(&self, state: &mut H) { self.expr.hash(state); @@ -1013,16 +1013,16 @@ where impl<'bind, 'parent, 'arena, T, A> IntoOrmExpression for CtxExpression<'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn into_orm_expression(self) -> OrmExpression { OrmExpression::Bound(self.expr) } } -fn bind_orm_context(executor: E, build: F) -> Result +fn bind_orm_context<'a, E, F>(executor: E, build: F) -> Result where - E: BindSource, + E: BindSource<'a>, F: for<'ctx, 'bind, 'parent, 'arena> FnOnce( &'ctx mut OrmContext< 'ctx, @@ -1030,20 +1030,22 @@ where 'parent, 'arena, E::Transaction, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { - static EMPTY_BIND_PARAMS: &[(&str, DataValue)] = &[]; - executor.execute(EMPTY_BIND_PARAMS, |binder, arena| { - let mut context = OrmContext { binder, arena }; - build(&mut context) + static EMPTY_BIND_PARAMS: &[(usize, LogicalType)] = &[]; + executor.execute(|state, tx| { + state.build_plan(EMPTY_BIND_PARAMS, tx, |binder, arena| { + let mut context = OrmContext { binder, arena }; + build(&mut context) + }) }) } -fn explain_orm_context(executor: E, build: F) -> Result +fn explain_orm_context<'a, E, F>(executor: E, build: F) -> Result where - E: BindSource, + E: BindSource<'a>, F: for<'ctx, 'bind, 'parent, 'arena> FnOnce( &'ctx mut OrmContext< 'ctx, @@ -1051,11 +1053,11 @@ where 'parent, 'arena, E::Transaction, - &'static [(&'static str, DataValue)], + &'static [(usize, LogicalType)], >, ) -> Result, { - static EMPTY_BIND_PARAMS: &[(&str, DataValue)] = &[]; + static EMPTY_BIND_PARAMS: &[(usize, LogicalType)] = &[]; executor.explain(EMPTY_BIND_PARAMS, |binder, arena| { let mut context = OrmContext { binder, arena }; build(&mut context) @@ -1070,7 +1072,7 @@ where pub struct OrmContext<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'ctx mut Binder<'bind, 'parent, T, A>, arena: &'ctx mut PlanArena<'arena>, @@ -1080,7 +1082,7 @@ where pub struct ExprBindScope<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'ctx mut Binder<'bind, 'parent, T, A>, arena: &'ctx mut PlanArena<'arena>, @@ -1089,7 +1091,7 @@ where pub struct UpdateBindScope<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { binder: &'ctx mut Binder<'bind, 'parent, T, A>, arena: &'ctx mut PlanArena<'arena>, @@ -1100,7 +1102,7 @@ where impl<'ctx, 'bind, 'parent, 'arena, T, A> OrmContext<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub fn from<'scope, M: Model>( &'scope mut self, @@ -1321,7 +1323,7 @@ where impl<'ctx, 'bind, 'parent, 'arena, T, A> ExprBindScope<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { fn handle(&self) -> ExprBindScopeHandle<'bind, 'parent, 'arena, T, A> { ExprBindScopeHandle::new(self) @@ -1826,7 +1828,7 @@ where impl<'ctx, 'bind, 'parent, 'arena, T, A> UpdateBindScope<'ctx, 'bind, 'parent, 'arena, T, A> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { pub fn set_value(&mut self, field: Field, value: D) -> Result<(), DatabaseError> where @@ -1922,7 +1924,7 @@ impl<'scope_ctx, 'bind, 'parent, 'arena, T, A, M> BindPlanFrom<'scope_ctx, 'bind, 'parent, 'arena, T, A, M> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, M: Model, { fn model_table_name(&self) -> Result { @@ -2375,7 +2377,7 @@ impl<'scope_ctx, 'bind, 'parent, 'arena, T, A, M> BindPlanSelectList<'scope_ctx, 'bind, 'parent, 'arena, T, A, M> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, M: Model, { fn expr_scope<'scope>(&'scope mut self) -> ExprBindScope<'scope, 'bind, 'parent, 'arena, T, A> { @@ -2544,7 +2546,7 @@ pub trait Projection: FromQueryRow { ) -> Result, DatabaseError> where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>; + A: AsRef<[(usize, LogicalType)]>; } fn orm_table_alias(source: &QuerySource) -> Option { @@ -2562,7 +2564,7 @@ fn bind_orm_source<'bind, 'parent, 'arena, T, A>( ) -> Result where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { let alias = orm_table_alias(&source); binder.bind_base_table_ref(join_type, source.table_name.as_str().into(), alias, arena) @@ -2576,7 +2578,7 @@ fn bind_orm_target_column<'bind, 'parent, 'arena, T, A>( ) -> Result where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { match binder.bind_column_ref_by_name(None, column_name, Some(source_name), arena)? { ScalarExpression::ColumnRef { column, .. } => Ok(column), @@ -2594,7 +2596,7 @@ fn bind_orm_insert_plan<'bind, 'parent, 'arena, T, A>( ) -> Result where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, { let table_name: TableName = table_name.into(); let input_schema = input_plan.output_schema(arena).clone(); @@ -2662,7 +2664,7 @@ fn bind_orm_insert_model<'bind, 'parent, 'arena, T, A, M>( ) -> Result where T: Transaction, - A: AsRef<[(&'static str, DataValue)]>, + A: AsRef<[(usize, LogicalType)]>, M: Model, { let table_name: TableName = M::table_name().into(); @@ -2691,10 +2693,10 @@ where )); } schema_ref.push(column); - row.push(value); + row.push(arena.alloc_expression(ScalarExpression::Constant(value))); } - binder.bind_insert_values(table_name, schema_ref, vec![row], false, true) + binder.bind_insert_values(table_name, schema_ref, row, 1, false, true) } fn describe_text_value(value: Option) -> String { @@ -3335,24 +3337,31 @@ fn extract_projected_tuple(tuple: &mut Tuple) -> Result(executor: E) -> Result<(), DatabaseError> { +fn orm_analyze<'a, E: BindSource<'a>, M: Model>(executor: E) -> Result<(), DatabaseError> { executor - .execute(&[], |binder, arena| { - binder.bind_analyze(M::table_name().into(), arena) + .execute(|state, tx| { + state.build_plan(&[], tx, |binder, arena| { + binder.bind_analyze(M::table_name().into(), arena) + }) })? .done() } -fn orm_insert(executor: E, model: &M) -> Result<(), DatabaseError> { +fn orm_insert<'a, E: BindSource<'a>, M: Model>( + executor: E, + model: &M, +) -> Result<(), DatabaseError> { let params = model.params(); executor - .execute(&[], |binder, arena| { - bind_orm_insert_model::<_, _, M>(binder, params, arena) + .execute(|state, tx| { + state.build_plan(&[], tx, |binder, arena| { + bind_orm_insert_model::<_, _, M>(binder, params, arena) + }) })? .done() } -fn orm_get( +fn orm_get<'a, E: BindSource<'a>, M: Model>( executor: E, key: &M::PrimaryKey, ) -> Result, DatabaseError> { @@ -3373,7 +3382,9 @@ fn orm_get( })?) } -fn orm_list(executor: E) -> Result, DatabaseError> { +fn orm_list<'a, E: BindSource<'a>, M: Model>( + executor: E, +) -> Result, DatabaseError> { Ok(bind_orm_context(executor, |ctx| { let plan: LogicalPlan = ctx.from::()?.finish()?; Ok(plan) @@ -3922,7 +3933,7 @@ mod tests { concat!( "Projection [upper_name] [Project => (Sort Option: Follow)] ", "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 ((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}" diff --git a/src/planner/arena.rs b/src/planner/arena.rs index dafd9e63..b96bb8e2 100644 --- a/src/planner/arena.rs +++ b/src/planner/arena.rs @@ -13,10 +13,12 @@ // limitations under the License. use crate::catalog::{ColumnCatalog, ColumnRef, TableName}; +use crate::errors::DatabaseError; use crate::expression::ScalarExpression; use crate::planner::LogicalPlan; use crate::types::index::{IndexMeta, IndexMetaRef}; use crate::types::tuple::Schema; +use crate::types::value::DataValue; use std::cell::UnsafeCell; use std::collections::HashSet; use std::fmt; @@ -40,7 +42,7 @@ struct TableArenaIndex { } struct TableArenaExpression { - expression: ScalarExpression, + expression: ArenaExpr, live: bool, } @@ -53,7 +55,7 @@ pub struct TableArenaCell { unsafe impl Send for TableArenaCell {} unsafe impl Sync for TableArenaCell {} -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct PlanArena<'a> { table_arena: &'a TableArenaCell, #[cfg(debug_assertions)] @@ -62,7 +64,7 @@ pub struct PlanArena<'a> { temp_table_id: usize, columns: Vec, indexes: Vec, - expressions: Vec, + expressions: Vec, plans: Vec, } @@ -92,7 +94,69 @@ impl fmt::Display for ExprRef { } } +#[derive(Debug, Clone)] +struct ArenaExpr { + expression: ScalarExpression, + param: Option, +} + +impl ArenaExpr { + fn new(expression: ScalarExpression) -> Self { + Self { + expression, + param: None, + } + } + + fn has_parameter(expression: &ScalarExpression) -> bool { + matches!(expression, ScalarExpression::Constant(value) if value.has_parameter()) + } +} + +/// Mutable expression access that invalidates a parameter slot when its expression changes. +pub struct ArenaExprMut<'a> { + expr: &'a mut ArenaExpr, +} + +impl std::ops::Deref for ArenaExprMut<'_> { + type Target = ScalarExpression; + + fn deref(&self) -> &Self::Target { + &self.expr.expression + } +} + +impl std::ops::DerefMut for ArenaExprMut<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + self.expr.param = None; + &mut self.expr.expression + } +} + pub trait MetaArena { + fn table_arena_cell<'a>(&self) -> &'a TableArenaCell + where + Self: 'a, + { + panic!("arena is not associated with a table catalog") + } + + fn expression_mut(&mut self, _expr: ExprRef) -> ArenaExprMut<'_> { + panic!("parent expressions are immutable") + } + + fn alloc_dummy(&mut self, name: &str) -> ColumnRef { + self.table_arena_cell().borrow().alloc_dummy(name) + } + + fn same_column(&self, left: ColumnRef, right: ColumnRef) -> bool { + self.column(left).summary() == self.column(right).summary() + } + + fn clone_column(&self, column: ColumnRef) -> ColumnCatalog { + self.column(column).clone() + } + fn alloc_column(&mut self, column: ColumnCatalog) -> ColumnRef; fn alloc_index(&mut self, index: IndexMeta) -> IndexMetaRef; @@ -121,6 +185,42 @@ pub trait MetaArena { fn find_index(&self, index: &IndexMeta) -> Option; } +impl MetaArena for Box { + fn table_arena_cell<'a>(&self) -> &'a TableArenaCell + where + Self: 'a, + { + (**self).table_arena_cell() + } + fn alloc_column(&mut self, column: ColumnCatalog) -> ColumnRef { + (**self).alloc_column(column) + } + fn alloc_index(&mut self, index: IndexMeta) -> IndexMetaRef { + (**self).alloc_index(index) + } + fn alloc_expression(&mut self, expr: ScalarExpression) -> ExprRef { + (**self).alloc_expression(expr) + } + fn expression_mut(&mut self, expr: ExprRef) -> ArenaExprMut<'_> { + (**self).expression_mut(expr) + } + fn column(&self, column: ColumnRef) -> &ColumnCatalog { + (**self).column(column) + } + fn index(&self, index: IndexMetaRef) -> &IndexMeta { + (**self).index(index) + } + fn expression(&self, expr: ExprRef) -> &ScalarExpression { + (**self).expression(expr) + } + fn find_column(&self, column: &ColumnCatalog) -> Option { + (**self).find_column(column) + } + fn find_index(&self, index: &IndexMeta) -> Option { + (**self).find_index(index) + } +} + const DUMMY_COLUMN_NAMES: [&str; DUMMY_COLUMN_COUNT] = [ "TABLE", "VIEW", @@ -344,6 +444,7 @@ impl MetaArena for TableArena { } fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + let expression = ArenaExpr::new(expression); if let Some((pos, slot)) = self .expressions .iter_mut() @@ -388,7 +489,7 @@ impl MetaArena for TableArena { fn expression(&self, expression: ExprRef) -> &ScalarExpression { let expression = &self.expressions[expression.pos()]; assert!(expression.live, "accessing recycled TableArena expression"); - &expression.expression + &expression.expression.expression } fn find_column(&self, column: &ColumnCatalog) -> Option { @@ -423,6 +524,50 @@ impl<'a> PlanArena<'a> { } } + fn arena_expression(&self, expression: ExprRef) -> &ArenaExpr { + self.assert_table_arena_unchanged(); + let table_arena = self.table_arena.borrow(); + let persistent_len = table_arena.expressions.len(); + if expression.pos() < persistent_len { + let slot = &table_arena.expressions[expression.pos()]; + assert!(slot.live, "accessing recycled TableArena expression"); + &slot.expression + } else { + &self.expressions[expression.pos() - persistent_len] + } + } + + pub(crate) fn expression_end(&self) -> usize { + self.table_arena.borrow().expressions.len() + self.expressions.len() + } + + pub(crate) fn parameter_expressions(&mut self) -> Vec { + let base = self.expression_end() - self.expressions.len(); + let mut parameters = Vec::new(); + for (offset, expr) in self.expressions.iter_mut().enumerate() { + expr.param = if ArenaExpr::has_parameter(&expr.expression) { + let slot = parameters.len(); + parameters.push(ExprRef::new(base + offset)); + Some(slot) + } else { + None + }; + } + parameters + } + + pub(crate) fn fill_parameters( + &mut self, + params: &[(usize, crate::types::value::DataValue)], + ) -> Result<(), crate::errors::DatabaseError> { + for expression in &mut self.expressions { + if let ScalarExpression::Constant(value) = &mut *(ArenaExprMut { expr: expression }) { + value.bind_parameters(params)?; + } + } + Ok(()) + } + pub(crate) fn table_arena_cell(&self) -> &'a TableArenaCell { self.table_arena } @@ -555,14 +700,16 @@ impl<'a> PlanArena<'a> { ::expression(self, expression_ref) } - pub(crate) fn expression_mut(&mut self, expression_ref: ExprRef) -> &mut ScalarExpression { + pub(crate) fn expression_mut(&mut self, expression_ref: ExprRef) -> ArenaExprMut<'_> { 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] + ArenaExprMut { + expr: &mut self.expressions[expression_ref.pos() - persistent_len], + } } pub fn column(&self, column: ColumnRef) -> &ColumnCatalog { @@ -576,6 +723,16 @@ impl<'a> PlanArena<'a> { } impl MetaArena for PlanArena<'_> { + fn table_arena_cell<'a>(&self) -> &'a TableArenaCell + where + Self: 'a, + { + self.table_arena + } + fn expression_mut(&mut self, expr: ExprRef) -> ArenaExprMut<'_> { + PlanArena::expression_mut(self, expr) + } + fn alloc_column(&mut self, column: ColumnCatalog) -> ColumnRef { self.assert_table_arena_unchanged(); self.allocated_columns_len += 1; @@ -604,7 +761,7 @@ impl MetaArena for PlanArena<'_> { 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); + self.expressions.push(ArenaExpr::new(expression)); ExprRef::new(pos) } @@ -633,14 +790,7 @@ impl MetaArena for PlanArena<'_> { } } 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] - } + &self.arena_expression(expression).expression } fn find_column(&self, column: &ColumnCatalog) -> Option { @@ -672,6 +822,94 @@ impl MetaArena for PlanArena<'_> { } } +/// Execution-local parameter values and temporary expressions over an immutable plan arena. +pub(crate) struct ParamArena<'a> { + parent: &'a PlanArena<'a>, + expressions: Vec, + parameter_count: usize, + parent_end: usize, +} + +impl<'a> ParamArena<'a> { + pub(crate) fn new( + parent: &'a PlanArena<'a>, + parameter_expressions: &[ExprRef], + params: &[(usize, DataValue)], + ) -> Result { + let mut expressions = Vec::with_capacity(parameter_expressions.len()); + for (slot, &expr_ref) in parameter_expressions.iter().enumerate() { + debug_assert_eq!(parent.arena_expression(expr_ref).param, Some(slot)); + let ScalarExpression::Constant(value) = parent.expression(expr_ref) else { + unreachable!("parameter expression must be a constant"); + }; + let mut value = value.clone(); + value.bind_parameters(params)?; + expressions.push(ArenaExpr::new(ScalarExpression::Constant(value))); + } + Ok(Self { + parent, + expressions, + parameter_count: parameter_expressions.len(), + parent_end: parent.expression_end(), + }) + } +} + +impl MetaArena for ParamArena<'_> { + fn table_arena_cell<'a>(&self) -> &'a TableArenaCell + where + Self: 'a, + { + self.parent.table_arena_cell() + } + fn column(&self, column: ColumnRef) -> &ColumnCatalog { + self.parent.column(column) + } + fn index(&self, index: IndexMetaRef) -> &IndexMeta { + self.parent.index(index) + } + fn find_column(&self, column: &ColumnCatalog) -> Option { + self.parent.find_column(column) + } + fn find_index(&self, index: &IndexMeta) -> Option { + self.parent.find_index(index) + } + fn expression(&self, expr: ExprRef) -> &ScalarExpression { + if expr.pos() < self.parent_end { + let parent = self.parent.arena_expression(expr); + return match parent.param { + Some(slot) => &self.expressions[slot].expression, + None => &parent.expression, + }; + } + &self.expressions[self.parameter_count + expr.pos() - self.parent_end].expression + } + fn expression_mut(&mut self, expr: ExprRef) -> ArenaExprMut<'_> { + let index = if expr.pos() < self.parent_end { + self.parent + .arena_expression(expr) + .param + .expect("parent expressions are immutable") + } else { + self.parameter_count + expr.pos() - self.parent_end + }; + ArenaExprMut { + expr: &mut self.expressions[index], + } + } + fn alloc_expression(&mut self, expression: ScalarExpression) -> ExprRef { + let id = ExprRef::new(self.parent_end + self.expressions.len() - self.parameter_count); + self.expressions.push(ArenaExpr::new(expression)); + id + } + fn alloc_column(&mut self, _: ColumnCatalog) -> ColumnRef { + unreachable!("ParamArena cannot allocate catalog columns") + } + fn alloc_index(&mut self, _: IndexMeta) -> IndexMetaRef { + unreachable!("ParamArena cannot allocate catalog indexes") + } +} + // GRCOV_EXCL_START #[cfg(test)] mod tests { @@ -707,6 +945,185 @@ mod tests { } } + #[test] + fn arena_expression_tracks_parameter_mutations() -> Result<(), DatabaseError> { + let root = TableArenaCell::default(); + let mut parent = PlanArena::new(&root); + let id = parent.alloc_expression(ScalarExpression::Constant(DataValue::Int32(1))); + assert_eq!(parent.arena_expression(id).param, None); + *parent.expression_mut(id) = + ScalarExpression::Constant(DataValue::Tuple(vec![DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + }])); + let parameter_expressions = parent.parameter_expressions(); + assert_eq!(parent.arena_expression(id).param, Some(0)); + { + let bound = + ParamArena::new(&parent, ¶meter_expressions, &[(1, DataValue::Int32(7))])?; + assert_eq!( + bound.expression(id), + &ScalarExpression::Constant(DataValue::Tuple(vec![DataValue::Int32(7)])) + ); + } + *parent.expression_mut(id) = ScalarExpression::Constant(DataValue::Int32(2)); + assert_eq!(parent.arena_expression(id).param, None); + let parameter_expressions = parent.parameter_expressions(); + let bound = ParamArena::new(&parent, ¶meter_expressions, &[])?; + assert!(std::ptr::eq(bound.expression(id), parent.expression(id))); + Ok(()) + } + + #[test] + fn arena_expression_updates_flag_on_unwind() { + let mut expression = ArenaExpr::new(ScalarExpression::Constant(DataValue::Int32(1))); + let result = catch_unwind(AssertUnwindSafe(|| { + let mut value = ArenaExprMut { + expr: &mut expression, + }; + *value = ScalarExpression::Constant(DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + }); + panic!("test mutation unwind"); + })); + assert!(result.is_err()); + assert!(ArenaExpr::has_parameter(&expression.expression)); + assert_eq!(expression.param, None); + } + + #[test] + fn param_arena_slots_are_dense_and_temporaries_follow() -> Result<(), DatabaseError> { + let root = TableArenaCell::default(); + let mut parent = PlanArena::new(&root); + let first = parent.alloc_expression(ScalarExpression::Constant(DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + })); + let middle = parent.alloc_expression(ScalarExpression::Constant(DataValue::Int32(99))); + let second = parent.alloc_expression(ScalarExpression::Constant(DataValue::Parameter { + id: 2, + ty: LogicalType::Integer, + })); + let parameters = parent.parameter_expressions(); + assert_eq!(parameters, [first, second]); + assert_eq!(parent.arena_expression(first).param, Some(0)); + assert_eq!(parent.arena_expression(middle).param, None); + assert_eq!(parent.arena_expression(second).param, Some(1)); + + let mut bound = ParamArena::new( + &parent, + ¶meters, + &[(1, DataValue::Int32(3)), (2, DataValue::Int32(8))], + )?; + assert_eq!( + bound.expression(first), + &ScalarExpression::Constant(DataValue::Int32(3)) + ); + assert_eq!( + bound.expression(second), + &ScalarExpression::Constant(DataValue::Int32(8)) + ); + assert!(std::ptr::eq( + bound.expression(middle), + parent.expression(middle) + )); + let temporary = bound.alloc_expression(ScalarExpression::Constant(DataValue::Int32(42))); + assert_eq!(temporary.pos(), parent.expression_end()); + assert_eq!( + bound.expression(temporary), + &ScalarExpression::Constant(DataValue::Int32(42)) + ); + *bound.expression_mut(second) = ScalarExpression::Constant(DataValue::Int32(9)); + *bound.expression_mut(temporary) = ScalarExpression::Constant(DataValue::Int32(43)); + let next = bound.alloc_expression(ScalarExpression::Constant(DataValue::Int32(44))); + assert_eq!(next.pos(), temporary.pos() + 1); + assert_eq!( + bound.expression(next), + &ScalarExpression::Constant(DataValue::Int32(44)) + ); + assert_eq!( + bound.expression(second), + &ScalarExpression::Constant(DataValue::Int32(9)) + ); + assert_eq!( + bound.expression(temporary), + &ScalarExpression::Constant(DataValue::Int32(43)) + ); + Ok(()) + } + + #[test] + fn param_arena_keeps_allocations_local() -> Result<(), DatabaseError> { + let root = TableArenaCell::default(); + let mut parent = PlanArena::new(&root); + let parameter = parent.alloc_expression(ScalarExpression::Constant(DataValue::Parameter { + id: 1, + ty: LogicalType::Integer, + })); + let constant = parent.alloc_expression(ScalarExpression::Constant(DataValue::Int32(20))); + let parent_column = parent.alloc_column(column("parent")); + let parent_index = parent.alloc_index(index_meta("parent")); + let parameter_expressions = parent.parameter_expressions(); + let mut first = + ParamArena::new(&parent, ¶meter_expressions, &[(1, DataValue::Int32(3))])?; + let second = ParamArena::new(&parent, ¶meter_expressions, &[(1, DataValue::Int32(8))])?; + assert!(std::ptr::eq( + first.expression(constant), + parent.expression(constant) + )); + assert_eq!( + first.expression(parameter), + &ScalarExpression::Constant(DataValue::Int32(3)) + ); + assert_eq!( + second.expression(parameter), + &ScalarExpression::Constant(DataValue::Int32(8)) + ); + assert!(catch_unwind(AssertUnwindSafe(|| { + let _ = first.expression_mut(constant); + })) + .is_err()); + *first.expression_mut(parameter) = ScalarExpression::Constant(DataValue::Int32(99)); + assert_eq!( + first.expression(parameter), + &ScalarExpression::Constant(DataValue::Int32(99)) + ); + assert_eq!( + second.expression(parameter), + &ScalarExpression::Constant(DataValue::Int32(8)) + ); + assert_eq!( + parent.expression(constant), + &ScalarExpression::Constant(DataValue::Int32(20)) + ); + assert_eq!(second.expression(constant), parent.expression(constant)); + let added = first.alloc_expression(ScalarExpression::Constant(DataValue::Int32(42))); + assert_eq!(added.pos(), parent.expression_end()); + assert_eq!( + first.expression(added), + &ScalarExpression::Constant(DataValue::Int32(42)) + ); + *first.expression_mut(added) = ScalarExpression::Constant(DataValue::Int32(43)); + assert_eq!( + first.expression(added), + &ScalarExpression::Constant(DataValue::Int32(43)) + ); + assert!(std::ptr::eq( + first.column(parent_column), + parent.column(parent_column) + )); + assert!(std::ptr::eq( + first.index(parent_index), + parent.index(parent_index) + )); + assert_eq!(first.find_column(&column("parent")), Some(parent_column)); + assert_eq!(first.find_index(&index_meta("parent")), Some(parent_index)); + assert!(first.find_column(&column("local")).is_none()); + assert!(first.find_index(&index_meta("local")).is_none()); + Ok(()) + } + #[test] fn table_arena_reuses_recycled_slot() { let arena = crate::planner::TableArenaCell::default(); diff --git a/src/planner/mod.rs b/src/planner/mod.rs index bcadaafd..2e95e613 100644 --- a/src/planner/mod.rs +++ b/src/planner/mod.rs @@ -28,13 +28,13 @@ use kite_sql_serde_macros::ReferenceSerialization; use std::fmt; use std::hash::{Hash, Hasher}; -pub(crate) use arena::PlanRef; pub use arena::{ExprRef, MetaArena, PlanArena, TableArena, TableArenaCell}; +pub(crate) use arena::{ParamArena, PlanRef}; pub(crate) trait Explain { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result; + fn fmt(&self, arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>) -> fmt::Result; - fn explain<'a, 'p>(&'a self, arena: &'a PlanArena<'p>) -> ExplainDisplay<'a, 'p, Self> + fn explain<'a, 'p>(&'a self, arena: &'a (dyn MetaArena + 'p)) -> ExplainDisplay<'a, 'p, Self> where Self: Sized, { @@ -44,7 +44,7 @@ pub(crate) trait Explain { pub(crate) struct ExplainDisplay<'a, 'p, T: ?Sized> { value: &'a T, - arena: &'a PlanArena<'p>, + arena: &'a (dyn MetaArena + 'p), } impl fmt::Display for ExplainDisplay<'_, '_, T> { @@ -56,7 +56,7 @@ impl fmt::Display for ExplainDisplay<'_, '_, T> { pub(crate) fn fmt_explain_list( values: &[T], separator: &str, - arena: &PlanArena<'_>, + arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>, ) -> fmt::Result { for (index, value) in values.iter().enumerate() { @@ -154,27 +154,10 @@ impl LogicalPlan { pub(crate) fn clone_plan( &self, - arena: &mut PlanArena<'_>, + arena: &mut (dyn MetaArena + '_), ) -> 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)?; + OperatorExprVisitorMut::new(&mut ExprCloner, arena).visit_plan(&mut plan)?; Ok(plan) } @@ -203,7 +186,7 @@ impl LogicalPlan { f: &mut F, ) -> Result<(), DatabaseError> where - A: MetaArena, + A: MetaArena + ?Sized, F: FnMut(&crate::catalog::ColumnRef) + ?Sized, { self.operator @@ -219,7 +202,7 @@ impl LogicalPlan { pub fn output_schema<'plan>( &'plan mut self, - arena: &mut PlanArena, + arena: &mut (dyn MetaArena + '_), ) -> &'plan crate::types::tuple::Schema { let LogicalPlan { operator, @@ -230,7 +213,7 @@ impl LogicalPlan { output_schema.get_or_insert_with(|| Self::compute_output_schema(operator, childrens, arena)) } - pub fn take_schema(&mut self, arena: &mut PlanArena) -> crate::types::tuple::Schema { + pub fn take_schema(&mut self, arena: &mut (dyn MetaArena + '_)) -> crate::types::tuple::Schema { let LogicalPlan { operator, childrens, @@ -245,7 +228,7 @@ impl LogicalPlan { fn compute_output_schema( operator: &mut Operator, childrens: &mut Childrens, - arena: &mut PlanArena, + arena: &mut (dyn MetaArena + '_), ) -> crate::types::tuple::Schema { match operator { Operator::Filter(_) @@ -357,7 +340,7 @@ impl LogicalPlan { } fn dummy_schema( - arena: &mut PlanArena, + arena: &mut (dyn MetaArena + '_), names: [&str; N], ) -> crate::types::tuple::Schema { names @@ -382,7 +365,7 @@ impl LogicalPlan { } } - pub fn explain(&self, arena: &mut PlanArena, indentation: usize) -> String { + pub fn explain(&self, arena: &mut (dyn MetaArena + '_), indentation: usize) -> String { format!( "{:indent$}{}", "", @@ -393,7 +376,7 @@ impl LogicalPlan { } impl Explain for LogicalPlan { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>) -> fmt::Result { self.operator.fmt(arena, f)?; if let Some(physical_option) = &self.physical_option { @@ -414,7 +397,7 @@ impl Clone for LogicalPlan { operator: self.operator.clone(), childrens: self.childrens.clone(), physical_option: self.physical_option.clone(), - output_schema: None, + output_schema: self.output_schema.clone(), } } } @@ -438,7 +421,7 @@ impl Hash for LogicalPlan { } impl crate::serdes::ReferenceSerialization for LogicalPlan { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -468,7 +451,7 @@ impl crate::serdes::ReferenceSerialization for LogicalPlan { ) } - fn decode( + fn decode( reader: &mut R, context: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &crate::serdes::ReferenceTables, @@ -668,3 +651,45 @@ mod tests { } } // GRCOV_EXCL_STOP + +#[cfg(test)] +pub(crate) mod test { + use crate::expression::ScalarExpression; + use crate::planner::{ExprRef, PlanArena}; + + pub(crate) trait PlanArenaTestExt { + fn alloc_expressions(&mut self, expressions: I) -> Vec + where + I: IntoIterator, + E: Into; + + fn alloc_expression_rows(&mut self, rows: &[R]) -> Vec + where + R: AsRef<[E]>, + E: Clone + Into; + } + + impl PlanArenaTestExt for PlanArena<'_> { + fn alloc_expressions(&mut self, expressions: I) -> Vec + where + I: IntoIterator, + E: Into, + { + expressions + .into_iter() + .map(|expression| self.alloc_expression(expression.into())) + .collect() + } + + fn alloc_expression_rows(&mut self, rows: &[R]) -> Vec + where + R: AsRef<[E]>, + E: Clone + Into, + { + rows.iter() + .flat_map(|row| row.as_ref().iter().cloned()) + .map(|expression| self.alloc_expression(expression.into())) + .collect() + } + } +} diff --git a/src/planner/operator/aggregate.rs b/src/planner/operator/aggregate.rs index 09efe78a..6457aeb9 100644 --- a/src/planner/operator/aggregate.rs +++ b/src/planner/operator/aggregate.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::planner::operator::Operator; -use crate::planner::{fmt_explain_list, Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, ExprRef, LogicalPlan}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -45,7 +46,11 @@ impl AggregateOperator { } impl Explain for AggregateOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { f.write_str("Aggregate [")?; fmt_explain_list(&self.agg_calls, ", ", arena, f)?; f.write_str("]")?; diff --git a/src/planner/operator/alter_table/change_column.rs b/src/planner/operator/alter_table/change_column.rs index f3bae4ce..6c8b2a33 100644 --- a/src/planner/operator/alter_table/change_column.rs +++ b/src/planner/operator/alter_table/change_column.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::catalog::TableName; -use crate::planner::{Explain, ExprRef, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{Explain, ExprRef}; use crate::types::LogicalType; use kite_sql_serde_macros::ReferenceSerialization; @@ -42,7 +43,11 @@ pub struct ChangeColumnOperator { } impl Explain for ChangeColumnOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { write!( f, "Change {} -> {}.{} ({}, ", diff --git a/src/planner/operator/analyze.rs b/src/planner/operator/analyze.rs index 314b4d92..101c08d7 100644 --- a/src/planner/operator/analyze.rs +++ b/src/planner/operator/analyze.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::catalog::TableName; -use crate::planner::{fmt_explain_list, Explain, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain}; use crate::types::index::IndexMetaRef; use kite_sql_serde_macros::ReferenceSerialization; @@ -25,7 +26,11 @@ pub struct AnalyzeOperator { } impl Explain for AnalyzeOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 2c28a6db..5e00ca92 100644 --- a/src/planner/operator/copy_from_file.rs +++ b/src/planner/operator/copy_from_file.rs @@ -14,6 +14,7 @@ use crate::binder::copy::ExtSource; use crate::catalog::TableName; +use crate::planner::MetaArena; use crate::planner::{fmt_explain_list, Explain, PlanArena}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; @@ -26,7 +27,11 @@ pub struct CopyFromFileOperator { } impl Explain for CopyFromFileOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 638926e0..c367e6eb 100644 --- a/src/planner/operator/create_index.rs +++ b/src/planner/operator/create_index.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::catalog::{ColumnRef, TableName}; -use crate::planner::{fmt_explain_list, Explain, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain}; use crate::types::index::IndexType; use kite_sql_serde_macros::ReferenceSerialization; @@ -28,7 +29,11 @@ pub struct CreateIndexOperator { } impl Explain for CreateIndexOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 8da2499f..2d76b8bd 100644 --- a/src/planner/operator/filter.rs +++ b/src/planner/operator/filter.rs @@ -12,7 +12,8 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan}; use kite_sql_serde_macros::ReferenceSerialization; use super::Operator; @@ -38,7 +39,11 @@ impl FilterOperator { } impl Explain for FilterOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { write!( f, "Filter {}, Is Having: {}", diff --git a/src/planner/operator/join.rs b/src/planner/operator/join.rs index 8580eae4..b734e900 100644 --- a/src/planner/operator/join.rs +++ b/src/planner/operator/join.rs @@ -13,7 +13,8 @@ // limitations under the License. use super::{Operator, PlanImpl}; -use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{Childrens, Explain, ExprRef, LogicalPlan}; use kite_sql_serde_macros::ReferenceSerialization; use std::fmt; use std::fmt::Formatter; @@ -82,13 +83,13 @@ impl JoinOperator { } impl Explain for JoinOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{} Join{}", self.join_type, self.on.explain(arena)) } } impl Explain for JoinCondition { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, arena: &(dyn MetaArena + '_), f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { JoinCondition::On { on, filter } => { if !on.is_empty() { diff --git a/src/planner/operator/mod.rs b/src/planner/operator/mod.rs index 68a0fbeb..478947ac 100644 --- a/src/planner/operator/mod.rs +++ b/src/planner/operator/mod.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -86,7 +88,7 @@ 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::{fmt_explain_list, Explain, ExprRef, MetaArena, PlanArena}; +use crate::planner::{fmt_explain_list, Explain, ExprRef, MetaArena}; use crate::types::index::{IndexInfo, IndexMetaRef}; use kite_sql_serde_macros::ReferenceSerialization; @@ -208,7 +210,11 @@ pub enum PlanImpl { } impl Explain for ColumnRef { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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()) @@ -219,7 +225,11 @@ impl Explain for ColumnRef { } impl Explain for IndexMetaRef { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { f.write_str(&arena.index(*self).name) } } @@ -229,7 +239,7 @@ macro_rules! impl_display_explain { $( $(#[$meta])* impl Explain for $ty { - fn fmt(&self, _arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt(&self, _arena: &(dyn MetaArena + '_), f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { std::fmt::Display::fmt(self, f) } } @@ -243,7 +253,6 @@ impl_display_explain!( ScalarSubqueryOperator, FunctionScanOperator, LimitOperator, - ValuesOperator, DescribeOperator, InsertOperator, DeleteOperator, @@ -260,7 +269,11 @@ impl_display_explain!( ); impl Explain for Operator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { match self { Operator::Dummy => f.write_str("Dummy"), Operator::Aggregate(op) => Explain::fmt(op, arena, f), @@ -308,12 +321,12 @@ impl Explain for Operator { } impl Operator { - pub fn visit_referenced_columns( + pub fn visit_referenced_columns( &self, arena: &A, f: &mut impl FnMut(&A, &ColumnRef) -> bool, ) -> Result { - struct ReferencedColumnVisitor<'a, A, F> { + struct ReferencedColumnVisitor<'a, A: ?Sized, F> { arena: &'a A, f: &'a mut F, keep_going: bool, @@ -321,7 +334,7 @@ impl Operator { impl ExprVisitor for ReferencedColumnVisitor<'_, A, F> where - A: MetaArena, + A: MetaArena + ?Sized, F: FnMut(&A, &ColumnRef) -> bool, { fn visit(&mut self, expr: ExprRef, arena: &A) -> Result<(), DatabaseError> { @@ -341,7 +354,7 @@ impl Operator { impl<'operator, A, F> OperatorVisitor<'operator> for ReferencedColumnVisitor<'_, A, F> where - A: MetaArena, + A: MetaArena + ?Sized, F: FnMut(&A, &ColumnRef) -> bool, { fn visit_aggregate( @@ -545,7 +558,7 @@ impl Operator { pub fn any_referenced_column( &self, - arena: &PlanArena, + arena: &(dyn MetaArena + '_), mut predicate: impl FnMut(&ColumnRef) -> bool, ) -> Result { let mut found = false; @@ -558,7 +571,7 @@ impl Operator { pub fn all_referenced_columns( &self, - arena: &PlanArena, + arena: &(dyn MetaArena + '_), mut predicate: impl FnMut(&ColumnRef) -> bool, ) -> Result { let mut all = true; @@ -571,7 +584,11 @@ impl Operator { } impl Explain for PlanImpl { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { match self { PlanImpl::Dummy => f.write_str("Dummy"), PlanImpl::SimpleAggregate => f.write_str("SimpleAggregate"), @@ -613,7 +630,11 @@ impl Explain for PlanImpl { } impl Explain for SortOption { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { match self { SortOption::OrderBy { fields, @@ -630,7 +651,11 @@ impl Explain for SortOption { } impl Explain for PhysicalOption { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { write!( f, "{} => (Sort Option: {})", @@ -657,6 +682,7 @@ mod tests { use crate::planner::operator::set_membership::SetMembershipKind; use crate::planner::operator::sort::SortField; use crate::planner::operator::values::ValuesOperator; + use crate::planner::test::PlanArenaTestExt; use crate::planner::ExprRef; use crate::planner::{Childrens, LogicalPlan, TableArenaCell}; use crate::types::index::{IndexInfo, IndexMeta, IndexMetaRef, IndexType}; @@ -805,10 +831,11 @@ mod tests { let mut arena = PlanArena::new(&table_arena); let left = column("left", &mut arena); let right = column("right", &mut arena); - let values = Operator::Values(ValuesOperator { - rows: vec![vec![DataValue::Int32(1), DataValue::Int32(2)]], - schema_ref: vec![left, right], - }); + let values = Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[vec![DataValue::Int32(1), DataValue::Int32(2)]]), + 1, + vec![left, right], + )); assert!(values.any_referenced_column(&arena, |column| *column == right)?); assert!( @@ -1209,14 +1236,15 @@ mod tests { "Describe users", ), ( - Operator::Values(ValuesOperator { - rows: vec![ - vec![DataValue::Int32(1), DataValue::Int32(2)], + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[ + vec![DataValue::Int32(1)], vec![DataValue::Int32(3)], - ], - schema_ref: vec![id], - }), - "Values [1, 2], [3], RowsLen: 2", + ]), + 2, + vec![id], + )), + "Values [1], [3], RowsLen: 2", ), ( Operator::Analyze(AnalyzeOperator { diff --git a/src/planner/operator/project.rs b/src/planner/operator/project.rs index 97003019..01d74216 100644 --- a/src/planner/operator/project.rs +++ b/src/planner/operator/project.rs @@ -12,7 +12,8 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::planner::{fmt_explain_list, Explain, ExprRef, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain, ExprRef}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -21,7 +22,11 @@ pub struct ProjectOperator { } impl Explain for ProjectOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 a84d3b34..1a420655 100644 --- a/src/planner/operator/recursive_cte.rs +++ b/src/planner/operator/recursive_cte.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::planner::operator::Operator; -use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; @@ -35,7 +36,11 @@ impl RecursiveCteOperator { } impl Explain for RecursiveCteOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { f.write_str("Recursive CTE: [")?; fmt_explain_list(&self.schema_ref, ", ", arena, f)?; f.write_str("]") @@ -48,7 +53,11 @@ pub struct RecursiveScanOperator { } impl Explain for RecursiveScanOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 e659dc6e..47b538c0 100644 --- a/src/planner/operator/set_membership.rs +++ b/src/planner/operator/set_membership.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::planner::operator::Operator; -use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; @@ -63,7 +64,11 @@ impl SetMembershipOperator { } impl Explain for SetMembershipOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 3de06a32..d930db0e 100644 --- a/src/planner/operator/sort.rs +++ b/src/planner/operator/sort.rs @@ -12,7 +12,8 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::planner::{fmt_explain_list, Explain, ExprRef, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain, ExprRef}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -64,14 +65,22 @@ pub struct SortOperator { } impl Explain for SortOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { f.write_str("Sort By ")?; fmt_explain_list(&self.sort_fields, ", ", arena, f) } } impl Explain for SortField { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { let direction = if self.asc { "Asc" } else { "Desc" }; let nulls = if self.nulls_first { "Nulls First" diff --git a/src/planner/operator/table_scan.rs b/src/planner/operator/table_scan.rs index d5d8c397..599ca376 100644 --- a/src/planner/operator/table_scan.rs +++ b/src/planner/operator/table_scan.rs @@ -18,7 +18,8 @@ use crate::errors::DatabaseError; use crate::expression::ScalarExpression; use crate::iter_ext::Itertools; use crate::planner::operator::sort::SortField; -use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan}; use crate::storage::Bounds; use crate::types::index::IndexInfo; use kite_sql_serde_macros::ReferenceSerialization; @@ -41,7 +42,7 @@ impl TableScanOperator { table_name: TableName, table_catalog: &TableCatalog, with_pk: bool, - arena: &mut PlanArena, + arena: &mut (dyn MetaArena + '_), ) -> Result { // Fill all Columns in TableCatalog by default let columns = table_catalog.columns().copied().collect_vec(); @@ -91,7 +92,11 @@ impl TableScanOperator { } impl Explain for TableScanOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { write!(f, "TableScan {} -> [", self.table_name)?; fmt_explain_list(&self.columns, ", ", arena, f)?; f.write_str("]")?; diff --git a/src/planner/operator/top_k.rs b/src/planner/operator/top_k.rs index 39080631..7eb00cd9 100644 --- a/src/planner/operator/top_k.rs +++ b/src/planner/operator/top_k.rs @@ -14,7 +14,8 @@ use super::Operator; use crate::planner::operator::sort::SortField; -use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -43,7 +44,11 @@ impl TopKOperator { } impl Explain for TopKOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { write!(f, "Top {}, ", self.limit)?; if let Some(offset) = self.offset { write!(f, "Offset {offset}, ")?; diff --git a/src/planner/operator/union.rs b/src/planner/operator/union.rs index 30370f55..95af467d 100644 --- a/src/planner/operator/union.rs +++ b/src/planner/operator/union.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::planner::operator::Operator; -use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Childrens, Explain, LogicalPlan}; use crate::types::tuple::Schema; use kite_sql_serde_macros::ReferenceSerialization; @@ -45,7 +46,11 @@ impl UnionOperator { } impl Explain for UnionOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 e4c2204e..d35b1ff3 100644 --- a/src/planner/operator/update.rs +++ b/src/planner/operator/update.rs @@ -13,7 +13,8 @@ // limitations under the License. use crate::catalog::{ColumnRef, TableName}; -use crate::planner::{Explain, ExprRef, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{Explain, ExprRef}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -23,7 +24,11 @@ pub struct UpdateOperator { } impl Explain for UpdateOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + 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 { diff --git a/src/planner/operator/values.rs b/src/planner/operator/values.rs index 91e1b51c..0ebc5f2e 100644 --- a/src/planner/operator/values.rs +++ b/src/planner/operator/values.rs @@ -12,32 +12,47 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::iter_ext::Itertools; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain, ExprRef}; use crate::types::tuple::Schema; -use crate::types::value::DataValue; use kite_sql_serde_macros::ReferenceSerialization; -use std::fmt; -use std::fmt::Formatter; +use std::fmt::{self, Formatter}; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] pub struct ValuesOperator { - pub rows: Vec>, + pub(crate) rows: Vec, + pub(crate) row_count: usize, pub schema_ref: Schema, } -impl fmt::Display for ValuesOperator { - fn fmt(&self, f: &mut Formatter) -> fmt::Result { - let columns = self - .rows - .iter() - .map(|row| { - let row_string = row.iter().map(|value| format!("{value}")).join(", "); - format!("[{row_string}]") - }) - .join(", "); - - write!(f, "Values {}, RowsLen: {}", columns, self.rows.len())?; +impl ValuesOperator { + pub fn new(rows: Vec, row_count: usize, schema_ref: Schema) -> Self { + assert_eq!( + Some(rows.len()), + row_count.checked_mul(schema_ref.len()), + "VALUES row width must match its schema" + ); + Self { + rows, + row_count, + schema_ref, + } + } +} - Ok(()) +impl Explain for ValuesOperator { + fn fmt(&self, arena: &(dyn MetaArena + '_), f: &mut Formatter) -> fmt::Result { + f.write_str("Values ")?; + let width = self.schema_ref.len(); + for i in 0..self.row_count { + let row = &self.rows[i * width..(i + 1) * width]; + if i != 0 { + f.write_str(", ")?; + } + f.write_str("[")?; + fmt_explain_list(row, ", ", arena, f)?; + f.write_str("]")?; + } + write!(f, ", RowsLen: {}", self.row_count) } } diff --git a/src/planner/operator/visitor.rs b/src/planner/operator/visitor.rs index c81f075f..c8aa0775 100644 --- a/src/planner/operator/visitor.rs +++ b/src/planner/operator/visitor.rs @@ -191,18 +191,27 @@ pub trait OperatorVisitor<'a>: Sized { } } -pub struct OperatorExprVisitor<'a, V, A> { +pub struct OperatorExprVisitor<'a, V, A: ?Sized> { visitor: &'a mut V, arena: &'a A, } -impl<'a, V, A> OperatorExprVisitor<'a, V, A> { +impl<'a, V, A: ?Sized> OperatorExprVisitor<'a, V, A> { pub fn new(visitor: &'a mut V, arena: &'a A) -> Self { Self { visitor, arena } } } -impl<'a, V: ExprVisitor, A: MetaArena> OperatorVisitor<'a> for OperatorExprVisitor<'_, V, A> { +impl<'a, V: ExprVisitor, A: MetaArena + ?Sized> OperatorVisitor<'a> + for OperatorExprVisitor<'_, V, A> +{ + fn visit_values(&mut self, op: &'a ValuesOperator) -> Result<(), DatabaseError> { + for expr in op.rows.iter() { + ExprVisitor::visit(self.visitor, *expr, self.arena)?; + } + Ok(()) + } + 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, self.arena)?; @@ -383,12 +392,12 @@ pub(crate) mod tests { }; use crate::expression::visitor::{walk_expr, ExprVisitor}; use crate::expression::window::{WindowFunction, WindowFunctionKind}; - use crate::expression::ScalarExpression; use crate::function::numbers::Numbers; use crate::planner::operator::alter_table::change_column::NotNullChange; use crate::planner::operator::join::{JoinOperator, JoinType}; use crate::planner::operator::mark_apply::MarkApplyOperator; use crate::planner::operator::set_membership::SetMembershipKind; + use crate::planner::test::PlanArenaTestExt; use crate::planner::ExprRef; use crate::planner::{Childrens, LogicalPlan}; use crate::types::index::{IndexInfo, IndexMetaRef, IndexType}; @@ -414,9 +423,7 @@ pub(crate) mod tests { 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 expressions = arena.alloc_expressions(0_i32..=18); let expr = |value: usize| expressions[value]; let column_ref = ColumnRef::new(0); let column = ColumnCatalog::new( @@ -480,10 +487,11 @@ pub(crate) mod tests { limit: 1, offset: None, }), - Operator::Values(ValuesOperator { - rows: vec![vec![DataValue::Int32(1)]], - schema_ref: vec![column_ref], - }), + Operator::Values(ValuesOperator::new( + arena.alloc_expression_rows(&[vec![DataValue::Int32(1)]]), + 1, + vec![column_ref], + )), Operator::Window(window::WindowOperator { sort_fields: vec![SortField::from(expr(17)), SortField::from(expr(18))], partition_by_len: 1, @@ -646,7 +654,7 @@ pub(crate) mod tests { for operator in &operators { visitor.visit_operator(operator)?; } - assert_eq!(counter.0, 20); + assert_eq!(counter.0, 21); // Includes the Values row expression. Ok(()) } diff --git a/src/planner/operator/visitor_mut.rs b/src/planner/operator/visitor_mut.rs index f8ac9c8a..d1d101ee 100644 --- a/src/planner/operator/visitor_mut.rs +++ b/src/planner/operator/visitor_mut.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +use crate::planner::PlanArena; // Copyright 2024 KipData/KiteSQL // // Licensed under the Apache License, Version 2.0 (the "License"); @@ -16,11 +18,45 @@ use super::alter_table::change_column::DefaultChange; use super::*; use crate::errors::DatabaseError; use crate::expression::visitor_mut::ExprVisitorMut; -use crate::planner::PlanArena; +use crate::planner::MetaArena; +use crate::planner::{Childrens, LogicalPlan}; pub trait OperatorVisitorMut<'a>: Sized { - fn visit_operator(&mut self, operator: &'a mut Operator) -> Result<(), DatabaseError> { - walk_mut_operator(self, operator) + fn visit_plan(&mut self, plan: &'a mut LogicalPlan) -> Result<(), DatabaseError> { + let LogicalPlan { + operator, + physical_option, + childrens, + .. + } = plan; + self.visit_operator(operator, physical_option.as_mut())?; + match childrens.as_mut() { + Childrens::Only(child) => self.visit_plan(child), + Childrens::Twins { left, right } => { + self.visit_plan(left)?; + self.visit_plan(right) + } + Childrens::None => Ok(()), + } + } + + fn visit_operator( + &mut self, + operator: &'a mut Operator, + physical_option: Option<&'a mut PhysicalOption>, + ) -> Result<(), DatabaseError> { + walk_mut_operator(self, operator)?; + if let Some(physical_option) = physical_option { + self.visit_physical_option(physical_option)?; + } + Ok(()) + } + + fn visit_physical_option( + &mut self, + _physical_option: &'a mut PhysicalOption, + ) -> Result<(), DatabaseError> { + Ok(()) } fn visit_dummy(&mut self) -> Result<(), DatabaseError> { @@ -214,16 +250,23 @@ pub trait OperatorVisitorMut<'a>: Sized { pub struct OperatorExprVisitorMut<'a, 'arena, V> { visitor: &'a mut V, - arena: &'a mut PlanArena<'arena>, + arena: &'a mut (dyn MetaArena + 'arena), } impl<'a, 'arena, V> OperatorExprVisitorMut<'a, 'arena, V> { - pub fn new(visitor: &'a mut V, arena: &'a mut PlanArena<'arena>) -> Self { + pub fn new(visitor: &'a mut V, arena: &'a mut (dyn MetaArena + 'arena)) -> Self { Self { visitor, arena } } } impl<'a, V: ExprVisitorMut> OperatorVisitorMut<'a> for OperatorExprVisitorMut<'_, '_, V> { + fn visit_values(&mut self, op: &'a mut ValuesOperator) -> Result<(), DatabaseError> { + for expr in op.rows.iter_mut() { + self.visitor.visit(expr, self.arena)?; + } + Ok(()) + } + 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, self.arena)?; @@ -414,7 +457,7 @@ mod tests { fn visit_constant( &mut self, value: &mut DataValue, - _arena: &mut PlanArena<'_>, + _arena: &mut (dyn MetaArena + '_), ) -> Result<(), DatabaseError> { if let DataValue::Int32(value) = value { *value += 1; @@ -424,6 +467,43 @@ mod tests { } } + #[test] + fn visits_plan_children_and_physical_options() -> Result<(), DatabaseError> { + struct Counter { + operators: usize, + physical_options: usize, + } + + impl<'a> OperatorVisitorMut<'a> for Counter { + fn visit_dummy(&mut self) -> Result<(), DatabaseError> { + self.operators += 1; + Ok(()) + } + + fn visit_physical_option( + &mut self, + _physical_option: &'a mut PhysicalOption, + ) -> Result<(), DatabaseError> { + self.physical_options += 1; + Ok(()) + } + } + + let mut child = LogicalPlan::new(Operator::Dummy, Childrens::None); + child.physical_option = Some(PhysicalOption::new(PlanImpl::Dummy, SortOption::None)); + let mut plan = LogicalPlan::new(Operator::Dummy, Childrens::Only(Box::new(child))); + plan.physical_option = Some(PhysicalOption::new(PlanImpl::Dummy, SortOption::None)); + + let mut counter = Counter { + operators: 0, + physical_options: 0, + }; + counter.visit_plan(&mut plan)?; + assert_eq!(counter.operators, 2); + assert_eq!(counter.physical_options, 2); + Ok(()) + } + #[test] fn dispatches_all_variants_and_mutates_expressions() -> Result<(), DatabaseError> { struct NoopVisitor; @@ -433,17 +513,17 @@ mod tests { let mut arena = PlanArena::new(&table_arena); let mut operators = all_operators(&mut arena)?; for operator in &mut operators { - NoopVisitor.visit_operator(operator)?; + NoopVisitor.visit_operator(operator, None)?; } let mut counter = IncrementConstants(0); { let mut visitor = OperatorExprVisitorMut::new(&mut counter, &mut arena); for operator in &mut operators { - visitor.visit_operator(operator)?; + visitor.visit_operator(operator, None)?; } } - assert_eq!(counter.0, 20); + assert_eq!(counter.0, 21); // Includes the Values row expression. Ok(()) } diff --git a/src/planner/operator/window.rs b/src/planner/operator/window.rs index 4483130b..0c200128 100644 --- a/src/planner/operator/window.rs +++ b/src/planner/operator/window.rs @@ -16,7 +16,8 @@ use crate::catalog::ColumnRef; use crate::expression::window::WindowFunction; use crate::planner::operator::sort::SortField; use crate::planner::operator::SortOption; -use crate::planner::{fmt_explain_list, Explain, PlanArena}; +use crate::planner::MetaArena; +use crate::planner::{fmt_explain_list, Explain}; use kite_sql_serde_macros::ReferenceSerialization; #[derive(Debug, PartialEq, Eq, Clone, Hash, ReferenceSerialization)] @@ -41,7 +42,11 @@ impl WindowOperator { } impl Explain for WindowOperator { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fn fmt( + &self, + arena: &(dyn MetaArena + '_), + f: &mut std::fmt::Formatter<'_>, + ) -> std::fmt::Result { let (partition_by, order_by) = self.sort_fields.split_at(self.partition_by_len); f.write_str("Window [")?; for (index, function) in self.functions.iter().enumerate() { diff --git a/src/python.rs b/src/python.rs index e83492e6..cca38f74 100644 --- a/src/python.rs +++ b/src/python.rs @@ -34,6 +34,11 @@ fn to_py_err(err: impl ToString) -> PyErr { #[allow(deprecated)] fn data_value_to_py(py: Python<'_>, value: &DataValue) -> PyResult { let object = match value { + DataValue::Parameter { id, .. } => { + return Err(PyRuntimeError::new_err(format!( + "unbound parameter ${id} reached Python output" + ))); + } DataValue::Null => py.None(), DataValue::Boolean(value) => value.into_py(py), DataValue::Float32(value) => value.0.into_py(py), @@ -52,7 +57,7 @@ fn data_value_to_py(py: Python<'_>, value: &DataValue) -> PyResult { | DataValue::Time32(_, _) | DataValue::Time64(_, _, _) | DataValue::Decimal(_) => value.to_string().into_py(py), - DataValue::Tuple(values, _is_upper) => { + DataValue::Tuple(values) => { let py_values = values .iter() .map(|value| data_value_to_py(py, value)) diff --git a/src/serdes/boolean.rs b/src/serdes/boolean.rs index d7447a27..7aaa6daf 100644 --- a/src/serdes/boolean.rs +++ b/src/serdes/boolean.rs @@ -13,12 +13,13 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; impl ReferenceSerialization for bool { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -28,7 +29,7 @@ impl ReferenceSerialization for bool { if *self { 1u8 } else { 0u8 }.encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/bound.rs b/src/serdes/bound.rs index 2faa8513..62a8c6d8 100644 --- a/src/serdes/bound.rs +++ b/src/serdes/bound.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; @@ -22,7 +23,7 @@ impl ReferenceSerialization for Bound where V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -48,7 +49,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/btree_map.rs b/src/serdes/btree_map.rs index 149114e2..5a51d497 100644 --- a/src/serdes/btree_map.rs +++ b/src/serdes/btree_map.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::collections::BTreeMap; @@ -23,7 +24,7 @@ where K: ReferenceSerialization + Ord, V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -39,7 +40,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/char.rs b/src/serdes/char.rs index e5049dae..a97511d8 100644 --- a/src/serdes/char.rs +++ b/src/serdes/char.rs @@ -13,12 +13,13 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; impl ReferenceSerialization for char { - fn encode( + fn encode( &self, writer: &mut W, _: bool, @@ -37,7 +38,7 @@ impl ReferenceSerialization for char { Ok(writer.write_all(&buf)?) } - fn decode( + fn decode( reader: &mut R, _: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, diff --git a/src/serdes/char_length_units.rs b/src/serdes/char_length_units.rs index d9f7725a..4d499b71 100644 --- a/src/serdes/char_length_units.rs +++ b/src/serdes/char_length_units.rs @@ -13,13 +13,14 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use crate::types::CharLengthUnits; use std::io::{Read, Write}; impl ReferenceSerialization for CharLengthUnits { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -35,7 +36,7 @@ impl ReferenceSerialization for CharLengthUnits { Ok(()) } - fn decode( + fn decode( reader: &mut R, _: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, diff --git a/src/serdes/column.rs b/src/serdes/column.rs index b5a7d2da..7ecb837b 100644 --- a/src/serdes/column.rs +++ b/src/serdes/column.rs @@ -21,7 +21,7 @@ use crate::types::ColumnId; use std::io::{Read, Write}; impl ReferenceSerialization for ColumnRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -33,7 +33,7 @@ impl ReferenceSerialization for ColumnRef { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -62,7 +62,7 @@ impl ReferenceSerialization for ColumnRef { } impl ReferenceSerialization for ColumnRelation { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -94,7 +94,7 @@ impl ReferenceSerialization for ColumnRelation { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/data_value.rs b/src/serdes/data_value.rs index 99fce987..e1c2bcde 100644 --- a/src/serdes/data_value.rs +++ b/src/serdes/data_value.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use crate::types::value::DataValue; @@ -44,7 +45,7 @@ const TAG_DECIMAL: u8 = 17; const TAG_TUPLE: u8 = 18; impl ReferenceSerialization for Utf8Type { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -63,7 +64,7 @@ impl ReferenceSerialization for Utf8Type { } } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -95,6 +96,9 @@ impl DataValue { writer: &mut W, ) -> Result<(), DatabaseError> { match self { + DataValue::Parameter { .. } => Err(DatabaseError::InvalidValue( + "unbound parameter cannot be serialized".to_string(), + )), DataValue::Null => write_u8(writer, TAG_NULL), DataValue::Boolean(value) => { write_u8(writer, TAG_BOOLEAN)?; @@ -171,13 +175,13 @@ impl DataValue { writer.write_all(&value.serialize())?; Ok(()) } - DataValue::Tuple(values, is_upper) => { + DataValue::Tuple(values) => { write_u8(writer, TAG_TUPLE)?; write_len(writer, values.len())?; for value in values { value.encode_reference_value(writer)?; } - write_bool(writer, *is_upper) + Ok(()) } } } @@ -229,7 +233,7 @@ impl DataValue { for _ in 0..len { values.push(DataValue::decode_reference_value(reader)?); } - Ok(DataValue::Tuple(values, read_bool(reader)?)) + Ok(DataValue::Tuple(values)) } tag => Err(DatabaseError::InvalidValue(format!( "invalid data value tag: {tag}" @@ -239,7 +243,7 @@ impl DataValue { } impl ReferenceSerialization for DataValue { - fn encode( + fn encode( &self, writer: &mut W, _: bool, @@ -249,7 +253,7 @@ impl ReferenceSerialization for DataValue { self.encode_reference_value(writer) } - fn decode( + fn decode( reader: &mut R, _: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, @@ -448,7 +452,7 @@ pub(crate) mod test { DataValue::Time64(78, 6, true), #[cfg(feature = "decimal")] DataValue::Decimal(Decimal::new(12345, 2)), - DataValue::Tuple(vec![DataValue::Null, DataValue::Int32(42)], false), + DataValue::Tuple(vec![DataValue::Null, DataValue::Int32(42)]), ]; let mut reference_tables = ReferenceTables::new(); @@ -475,6 +479,18 @@ pub(crate) mod test { Ok(()) } + #[test] + fn unbound_parameter_cannot_be_serialized() { + let parameter = DataValue::Parameter { + id: 1, + ty: crate::types::LogicalType::Integer, + }; + let err = parameter + .encode_reference_value(&mut Vec::new()) + .unwrap_err(); + assert!(err.to_string().contains("unbound parameter")); + } + #[test] fn invalid_tags_and_truncated_values_are_rejected() { assert_invalid_value(&[u8::MAX], "invalid data value tag"); @@ -503,9 +519,8 @@ pub(crate) mod test { assert_invalid_value(&time64, "invalid bool value"); let mut tuple = vec![TAG_TUPLE]; - tuple.extend(0u32.to_le_bytes()); - tuple.push(2); - assert_invalid_value(&tuple, "invalid bool value"); + tuple.extend(1u32.to_le_bytes()); + assert_invalid_value(&tuple, "failed to fill whole buffer"); } #[test] diff --git a/src/serdes/evaluator.rs b/src/serdes/evaluator.rs index 4f47af14..5fd8a466 100644 --- a/src/serdes/evaluator.rs +++ b/src/serdes/evaluator.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceDecodeContext, ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use crate::types::evaluator::{ @@ -22,7 +23,7 @@ use crate::types::evaluator::{ use std::io::{Read, Write}; impl ReferenceSerialization for BinaryEvaluatorParams { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -38,7 +39,7 @@ impl ReferenceSerialization for BinaryEvaluatorParams { } } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -57,7 +58,7 @@ impl ReferenceSerialization for BinaryEvaluatorParams { } impl ReferenceSerialization for CastEvaluatorParams { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -94,7 +95,7 @@ impl ReferenceSerialization for CastEvaluatorParams { } } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -148,7 +149,7 @@ impl ReferenceSerialization for CastEvaluatorParams { } impl ReferenceSerialization for UnaryEvaluatorRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -158,7 +159,7 @@ impl ReferenceSerialization for UnaryEvaluatorRef { self.pos.encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -174,7 +175,7 @@ impl ReferenceSerialization for UnaryEvaluatorRef { } impl ReferenceSerialization for BinaryEvaluatorRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -187,7 +188,7 @@ impl ReferenceSerialization for BinaryEvaluatorRef { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -201,7 +202,7 @@ impl ReferenceSerialization for BinaryEvaluatorRef { } impl ReferenceSerialization for CastEvaluatorRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -214,7 +215,7 @@ impl ReferenceSerialization for CastEvaluatorRef { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/expression.rs b/src/serdes/expression.rs index c516158c..c4ead6d9 100644 --- a/src/serdes/expression.rs +++ b/src/serdes/expression.rs @@ -20,7 +20,7 @@ use crate::storage::Transaction; use std::io::{Read, Write}; impl ReferenceSerialization for ExprRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -32,7 +32,7 @@ impl ReferenceSerialization for ExprRef { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/function.rs b/src/serdes/function.rs index 1627af02..d528212a 100644 --- a/src/serdes/function.rs +++ b/src/serdes/function.rs @@ -16,12 +16,13 @@ use crate::errors::DatabaseError; use crate::expression::function::scala::ArcScalarFunctionImpl; use crate::expression::function::table::ArcTableFunctionImpl; use crate::expression::function::FunctionSummary; +use crate::planner::MetaArena; use crate::serdes::{ReferenceDecodeContext, ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; impl ReferenceSerialization for ArcScalarFunctionImpl { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -32,7 +33,7 @@ impl ReferenceSerialization for ArcScalarFunctionImpl { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -57,7 +58,7 @@ impl ReferenceSerialization for ArcScalarFunctionImpl { } impl ReferenceSerialization for ArcTableFunctionImpl { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -68,7 +69,7 @@ impl ReferenceSerialization for ArcTableFunctionImpl { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/hasher.rs b/src/serdes/hasher.rs index ae4bed8d..737fbbfa 100644 --- a/src/serdes/hasher.rs +++ b/src/serdes/hasher.rs @@ -13,12 +13,13 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::stable_hash::StableHasher; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; impl ReferenceSerialization for StableHasher { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -30,7 +31,7 @@ impl ReferenceSerialization for StableHasher { key1.encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/index.rs b/src/serdes/index.rs index e676da14..9aef83a9 100644 --- a/src/serdes/index.rs +++ b/src/serdes/index.rs @@ -20,7 +20,7 @@ use crate::types::index::{IndexMeta, IndexMetaRef}; use std::io::{Read, Write}; impl ReferenceSerialization for IndexMetaRef { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -32,7 +32,7 @@ impl ReferenceSerialization for IndexMetaRef { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/mod.rs b/src/serdes/mod.rs index eeddd439..48030e73 100644 --- a/src/serdes/mod.rs +++ b/src/serdes/mod.rs @@ -45,7 +45,7 @@ use std::io; use std::io::{Read, Write}; pub trait ReferenceSerialization { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -53,7 +53,7 @@ pub trait ReferenceSerialization { arena: &A, ) -> Result<(), DatabaseError>; - fn decode( + fn decode( reader: &mut R, context: Option<&ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/num.rs b/src/serdes/num.rs index 264040cb..a22db8da 100644 --- a/src/serdes/num.rs +++ b/src/serdes/num.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::Read; @@ -23,7 +24,7 @@ use std::mem::size_of; macro_rules! implement_num_serialization { ($struct_name:ident) => { impl ReferenceSerialization for $struct_name { - fn encode( + fn encode( &self, writer: &mut W, _: bool, @@ -35,7 +36,7 @@ macro_rules! implement_num_serialization { Ok(()) } - fn decode( + fn decode( reader: &mut R, _: Option<&$crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, @@ -63,7 +64,7 @@ implement_num_serialization!(f32); implement_num_serialization!(f64); impl ReferenceSerialization for usize { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -73,7 +74,7 @@ impl ReferenceSerialization for usize { (*self as u32).encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -86,6 +87,7 @@ impl ReferenceSerialization for usize { #[cfg(all(test, not(target_arch = "wasm32")))] pub(crate) mod test { use crate::errors::DatabaseError; + use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::rocksdb::RocksTransaction; use std::io::{Cursor, Seek, SeekFrom}; diff --git a/src/serdes/option.rs b/src/serdes/option.rs index 60b57991..82e2427f 100644 --- a/src/serdes/option.rs +++ b/src/serdes/option.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; @@ -21,7 +22,7 @@ impl ReferenceSerialization for Option where V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -39,7 +40,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/pair.rs b/src/serdes/pair.rs index f9c9b904..3fd86c02 100644 --- a/src/serdes/pair.rs +++ b/src/serdes/pair.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; @@ -22,7 +23,7 @@ where A: ReferenceSerialization, B: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -36,7 +37,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/path_buf.rs b/src/serdes/path_buf.rs index 9ecb6117..072dac6f 100644 --- a/src/serdes/path_buf.rs +++ b/src/serdes/path_buf.rs @@ -13,13 +13,14 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::path::PathBuf; use std::str::FromStr; impl ReferenceSerialization for PathBuf { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -32,7 +33,7 @@ impl ReferenceSerialization for PathBuf { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/phantom.rs b/src/serdes/phantom.rs index 40dad173..33e01b8e 100644 --- a/src/serdes/phantom.rs +++ b/src/serdes/phantom.rs @@ -13,13 +13,14 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; use std::marker::PhantomData; impl ReferenceSerialization for PhantomData { - fn encode( + fn encode( &self, _: &mut W, _: bool, @@ -29,7 +30,7 @@ impl ReferenceSerialization for PhantomData { Ok(()) } - fn decode( + fn decode( _: &mut R, _: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, diff --git a/src/serdes/ptr.rs b/src/serdes/ptr.rs index 7f0a5c69..e5301158 100644 --- a/src/serdes/ptr.rs +++ b/src/serdes/ptr.rs @@ -25,7 +25,7 @@ macro_rules! implement_ptr_serialization { where V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -36,7 +36,7 @@ macro_rules! implement_ptr_serialization { .encode(writer, is_direct, reference_tables, arena) } - fn decode( + fn decode( reader: &mut R, drive: Option<&$crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/slice.rs b/src/serdes/slice.rs index 97dcd019..1d4f567b 100644 --- a/src/serdes/slice.rs +++ b/src/serdes/slice.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; @@ -21,7 +22,7 @@ impl ReferenceSerialization for [V; 2] where V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -34,7 +35,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/string.rs b/src/serdes/string.rs index e5539d3d..378f950f 100644 --- a/src/serdes/string.rs +++ b/src/serdes/string.rs @@ -13,13 +13,14 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; use std::sync::Arc; impl ReferenceSerialization for String { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -33,7 +34,7 @@ impl ReferenceSerialization for String { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, @@ -48,7 +49,7 @@ impl ReferenceSerialization for String { } impl ReferenceSerialization for Arc { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -62,7 +63,7 @@ impl ReferenceSerialization for Arc { Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/serdes/trim.rs b/src/serdes/trim.rs index 42ca63fc..afbb42e1 100644 --- a/src/serdes/trim.rs +++ b/src/serdes/trim.rs @@ -14,12 +14,13 @@ use crate::errors::DatabaseError; use crate::expression::TrimWhereField; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; impl ReferenceSerialization for TrimWhereField { - fn encode( + fn encode( &self, writer: &mut W, _: bool, @@ -36,7 +37,7 @@ impl ReferenceSerialization for TrimWhereField { Ok(()) } - fn decode( + fn decode( reader: &mut R, _: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, _: &ReferenceTables, diff --git a/src/serdes/vec.rs b/src/serdes/vec.rs index 6e3a318b..5cbabd6c 100644 --- a/src/serdes/vec.rs +++ b/src/serdes/vec.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::errors::DatabaseError; +use crate::planner::MetaArena; use crate::serdes::{ReferenceSerialization, ReferenceTables}; use crate::storage::Transaction; use std::io::{Read, Write}; @@ -21,7 +22,7 @@ impl ReferenceSerialization for Vec where V: ReferenceSerialization, { - fn encode( + fn encode( &self, writer: &mut W, is_direct: bool, @@ -36,7 +37,7 @@ where Ok(()) } - fn decode( + fn decode( reader: &mut R, drive: Option<&crate::serdes::ReferenceDecodeContext<'_, T>>, reference_tables: &ReferenceTables, diff --git a/src/storage/mod.rs b/src/storage/mod.rs index 4fb48959..e9c31590 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -31,13 +31,13 @@ use crate::optimizer::core::cm_sketch::{ }; use crate::optimizer::core::statistics_meta::StatisticsMeta; use crate::planner::operator::alter_table::change_column::{DefaultChange, NotNullChange}; -use crate::planner::{MetaArena, PlanArena, TableArenaCell}; +use crate::planner::{MetaArena, TableArenaCell}; use crate::serdes::ReferenceTables; use crate::storage::table_codec::{Bytes, StatisticsCodecType, TableCodec, BOUND_MAX_TAG}; use crate::types::index::{Index, IndexId, IndexMeta, IndexMetaRef, IndexType}; use crate::types::serialize::TupleValueSerializableImpl; use crate::types::tuple::{Tuple, TupleId}; -use crate::types::value::{DataValue, TupleMappingRef}; +use crate::types::value::{DataValue, OrderedMapping, TupleMappingRef}; use crate::types::{ColumnId, LogicalType}; use std::borrow::{Borrow, Cow}; use std::collections::{Bound, HashMap}; @@ -76,7 +76,7 @@ impl Display for TransactionIsolationLevel { pub(crate) fn index_value_type( table: &TableCatalog, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), column_ids: &[ColumnId], ) -> Result { let mut value_types = Vec::with_capacity(column_ids.len()); @@ -177,7 +177,7 @@ pub trait Transaction: Sized { fn read<'a>( &'a self, table_codec: &mut TableCodec, - arena: &PlanArena, + arena: &(dyn MetaArena + '_), table_cache: &TableCache, table_name: TableName, bounds: Bounds, @@ -208,7 +208,7 @@ pub trait Transaction: Sized { fn read_by_index<'a, R>( &'a self, table_cache: &TableCache, - arena: &PlanArena<'a>, + arena: &(dyn MetaArena + 'a), table_name: TableName, (offset_option, limit_option): Bounds, columns: Vec, @@ -277,7 +277,7 @@ pub trait Transaction: Sized { fn create_deserializers( columns: &[ColumnRef], table: &TableCatalog, - arena: &PlanArena, + arena: &(dyn MetaArena + '_), with_pk: bool, ) -> Vec { let mut pk_len = if with_pk { @@ -319,7 +319,7 @@ pub trait Transaction: Sized { fn add_index_meta( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: &TableName, index_name: String, column_ids: Vec, @@ -422,7 +422,7 @@ pub trait Transaction: Sized { fn rewrite_table_metadata( &mut self, table_codec: &mut TableCodec, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), table: &TableCatalog, ) -> Result<(), DatabaseError> { let table_name = table.name().clone(); @@ -454,7 +454,7 @@ pub trait Transaction: Sized { fn change_column( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: &TableName, old_column_name: &str, new_column_name: &str, @@ -548,7 +548,7 @@ pub trait Transaction: Sized { fn add_column( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: &TableName, column: &ColumnCatalog, if_not_exists: bool, @@ -597,7 +597,7 @@ pub trait Transaction: Sized { fn drop_column( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: &TableName, column_name: &str, ) -> Result { @@ -637,7 +637,7 @@ pub trait Transaction: Sized { fn create_view( &mut self, table_codec: &mut TableCodec, - arena: &PlanArena, + arena: &(dyn MetaArena + '_), view: View, or_replace: bool, ) -> Result { @@ -656,7 +656,7 @@ pub trait Transaction: Sized { fn create_table( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: TableName, columns: Vec, if_not_exists: bool, @@ -740,7 +740,7 @@ pub trait Transaction: Sized { fn drop_index( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: TableName, index_name: &str, if_exists: bool, @@ -790,7 +790,7 @@ pub trait Transaction: Sized { fn drop_table( &mut self, table_codec: &mut TableCodec, - plan_arena: &mut PlanArena, + plan_arena: &mut (dyn MetaArena + '_), table_name: TableName, if_exists: bool, ) -> Result { @@ -908,7 +908,7 @@ pub trait Transaction: Sized { fn load_table( &self, table_codec: &mut TableCodec, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), table_name: TableName, ) -> Result, DatabaseError> { self.table_collect(table_codec, &table_name, arena)? @@ -934,7 +934,7 @@ pub trait Transaction: Sized { table_codec: &mut TableCodec, table_name: &TableName, statistics_meta: StatisticsMeta, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { let index_id = statistics_meta.index_id(); let (root, buckets, cm_sketch, top_n) = statistics_meta.into_parts(); @@ -991,7 +991,7 @@ pub trait Transaction: Sized { table_codec: &mut TableCodec, table_name: &str, index_id: IndexId, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result, DatabaseError> { table_codec.with_statistics_index_bound(table_name, index_id, |min, max| { let mut iter = self.range(Bound::Included(min), Bound::Included(max))?; @@ -1060,7 +1060,7 @@ pub trait Transaction: Sized { &self, table_codec: &mut TableCodec, table_name: &TableName, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result, Vec)>, DatabaseError> { table_codec.with_table_bound(table_name, |table_min, table_max| { let mut column_iter = @@ -1092,7 +1092,7 @@ pub trait Transaction: Sized { fn create_index_meta_from_column( &mut self, table_codec: &mut TableCodec, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), table: &mut TableCatalog, ) -> Result<(), DatabaseError> { let table_name = table.name.clone(); @@ -1154,6 +1154,7 @@ pub trait Transaction: Sized { fn remove(&mut self, key: &[u8]) -> Result<(), DatabaseError>; + // TODO: Support reverse range iteration (upper to lower bound) for descending index scans. fn range<'txn, 'key>( &'txn self, min: Bound<&'key [u8]>, @@ -1244,6 +1245,14 @@ fn encode_bound_key(buffer: &mut Bytes, key: &[u8], is_upper: bool) { } } +#[inline] +fn index_values(value: &DataValue) -> &[DataValue] { + match value { + DataValue::Tuple(values) => values, + value => std::slice::from_ref(value), + } +} + #[inline] fn encode_bound<'a>( table_codec: &mut TableCodec, @@ -1538,12 +1547,11 @@ impl IndexImpl for PrimaryKeyIndexImpl { table_codec: &mut TableCodec, params: &IndexImplParams, value: &DataValue, - _: bool, + is_upper: bool, out: &mut Bytes, ) -> Result<(), DatabaseError> { table_codec.with_tuple_unchecked(params.table_name.as_ref(), value, None, |key, _| { - out.clear(); - out.extend_from_slice(key); + encode_bound_key(out, key, is_upper); Ok(()) }) } @@ -1585,7 +1593,11 @@ impl IndexImpl for UniqueIndexImpl { _: &mut Bytes, _: &mut Bytes, ) -> Result, DatabaseError> { - let index = Index::new(params.index_meta().id, value, IndexType::Unique); + let index = Index::new( + params.index_meta().id, + index_values(value), + IndexType::Unique, + ); let Some(bytes) = table_codec.with_index(params.table_name.as_ref(), &index, None, |key, _| { params.tx.get_borrowed(key) @@ -1609,7 +1621,11 @@ impl IndexImpl for UniqueIndexImpl { _: bool, out: &mut Bytes, ) -> Result<(), DatabaseError> { - let index = Index::new(params.index_meta().id, value, IndexType::Unique); + let index = Index::new( + params.index_meta().id, + index_values(value), + IndexType::Unique, + ); table_codec.with_index(params.table_name.as_ref(), &index, None, |key, _| { out.clear(); @@ -1659,7 +1675,11 @@ impl IndexImpl for NormalIndexImpl { is_upper: bool, out: &mut Bytes, ) -> Result<(), DatabaseError> { - let index = Index::new(params.index_meta().id, value, IndexType::Normal); + let index = Index::new( + params.index_meta().id, + index_values(value), + IndexType::Normal, + ); table_codec.with_index(params.table_name.as_ref(), &index, None, |key, _| { encode_bound_key(out, key, is_upper); Ok(()) @@ -1707,7 +1727,11 @@ impl IndexImpl for CompositeIndexImpl { is_upper: bool, out: &mut Bytes, ) -> Result<(), DatabaseError> { - let index = Index::new(params.index_meta().id, value, IndexType::Composite); + let index = Index::new( + params.index_meta().id, + index_values(value), + IndexType::Composite, + ); table_codec.with_index(params.table_name.as_ref(), &index, None, |key, _| { encode_bound_key(out, key, is_upper); Ok(()) @@ -1724,11 +1748,12 @@ impl IndexImpl for CoveredIndexImpl { value: &[u8], params: &IndexImplParams, ) -> Result<(), DatabaseError> { - let mapping = params - .cover_mapping - .as_ref() - .map(|mapping| mapping.as_ref()); - let key = TableCodec::decode_index_key(key, params.value_ty(), mapping)?; + let key = match ¶ms.cover_mapping { + Some(mapping) => { + TableCodec::decode_index_key(key, params.value_ty(), &mapping.as_ref())? + } + None => TableCodec::decode_index_key(key, params.value_ty(), &OrderedMapping)?, + }; let tuple_id = if params.with_pk { Some(TableCodec::decode_index(value)?) @@ -1736,10 +1761,7 @@ impl IndexImpl for CoveredIndexImpl { None }; tuple.pk = tuple_id; - tuple.values = match key { - DataValue::Tuple(vals, _) => vals, - v => vec![v], - }; + tuple.values = key; Ok(()) } @@ -1771,7 +1793,11 @@ impl IndexImpl for CoveredIndexImpl { is_upper: bool, out: &mut Bytes, ) -> Result<(), DatabaseError> { - let index = Index::new(params.index_meta().id, value, params.index_meta().ty); + let index = Index::new( + params.index_meta().id, + index_values(value), + params.index_meta().ty, + ); table_codec.with_index(params.table_name.as_ref(), &index, None, |key, _| { encode_bound_key(out, key, is_upper); Ok(()) @@ -2062,7 +2088,7 @@ pub struct TableIter<'a, T: Transaction + 'a> { impl TableIter<'_, T> { pub fn try_next( &mut self, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result, DatabaseError> { let Some((_, value)) = self.iter.try_next()? else { return Ok(None); @@ -2576,15 +2602,27 @@ mod test { let indexes = [ ( Arc::new(DataValue::Int32(0)), - Index::new(1, &tuples[0].values[2], IndexType::Normal), + Index::new( + 1, + std::slice::from_ref(&tuples[0].values[2]), + IndexType::Normal, + ), ), ( Arc::new(DataValue::Int32(1)), - Index::new(1, &tuples[1].values[2], IndexType::Normal), + Index::new( + 1, + std::slice::from_ref(&tuples[1].values[2]), + IndexType::Normal, + ), ), ( Arc::new(DataValue::Int32(2)), - Index::new(1, &tuples[2].values[2], IndexType::Normal), + Index::new( + 1, + std::slice::from_ref(&tuples[2].values[2]), + IndexType::Normal, + ), ), ]; for (tuple_id, index) in indexes.iter().cloned() { @@ -2710,7 +2748,11 @@ mod test { setup_tx.add_index( &mut table_codec, "t1", - Index::new(index_id, &initial_tuple.values[2], IndexType::Normal), + Index::new( + index_id, + std::slice::from_ref(&initial_tuple.values[2]), + IndexType::Normal, + ), initial_tuple.pk.as_ref().unwrap(), )?; setup_tx.append_tuple(&mut table_codec, "t1", &initial_tuple, &serializers, false)?; @@ -2748,13 +2790,21 @@ mod test { writer_tx.del_index( &mut table_codec, "t1", - &Index::new(index_id, &initial_tuple.values[2], IndexType::Normal), + &Index::new( + index_id, + std::slice::from_ref(&initial_tuple.values[2]), + IndexType::Normal, + ), initial_tuple.pk.as_ref().unwrap(), )?; writer_tx.add_index( &mut table_codec, "t1", - Index::new(index_id, &updated_tuple.values[2], IndexType::Normal), + Index::new( + index_id, + std::slice::from_ref(&updated_tuple.values[2]), + IndexType::Normal, + ), updated_tuple.pk.as_ref().unwrap(), )?; writer_tx.append_tuple(&mut table_codec, "t1", &updated_tuple, &serializers, true)?; diff --git a/src/storage/table_codec.rs b/src/storage/table_codec.rs index 1c0145e5..d1d485d7 100644 --- a/src/storage/table_codec.rs +++ b/src/storage/table_codec.rs @@ -27,7 +27,7 @@ use crate::storage::{TableCache, Transaction}; use crate::types::index::{Index, IndexId, IndexMeta, IndexType, INDEX_ID_LEN}; use crate::types::serialize::TupleValueSerializableImpl; use crate::types::tuple::{Tuple, TupleId}; -use crate::types::value::{DataValue, TupleMappingRef}; +use crate::types::value::{DataValue, IndexKeyMapping}; use crate::types::LogicalType; use std::borrow::Borrow; use std::hash::{Hash, Hasher}; @@ -151,7 +151,7 @@ impl TableCodec { return Err(DatabaseError::not_null_column("primary key")); } - if let DataValue::Tuple(values, _) = &value { + if let DataValue::Tuple(values) = &value { for value in values { Self::check_primary_key(value, indentation + 1)? } @@ -307,7 +307,7 @@ impl TableCodec { table_name: &str, index_id: IndexId, index_meta: Option<&IndexMeta>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -401,7 +401,9 @@ impl TableCodec { lower.push(BOUND_MIN_TAG); lower.extend_from_slice(&index.id.to_le_bytes()); lower.push(BOUND_MIN_TAG); - index.value.memcomparable_encode(lower)?; + for value in index.values { + value.memcomparable_encode(lower)?; + } if let Some(tuple_id) = tuple_id { if matches!(index.ty, IndexType::Normal | IndexType::Composite) { @@ -422,7 +424,7 @@ impl TableCodec { &mut self, col: &ColumnCatalog, encode_value: bool, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { if let ColumnRelation::Table { @@ -532,7 +534,7 @@ impl TableCodec { table_name: &str, index_id: IndexId, statistics_meta: Option<&StatisticsMetaRoot>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -556,7 +558,7 @@ impl TableCodec { table_name: &str, index_id: IndexId, sketch_meta: Option<&CountMinSketchMeta>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -581,7 +583,7 @@ impl TableCodec { index_id: IndexId, sketch_page: &CountMinSketchPage, encode_value: bool, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -610,7 +612,7 @@ impl TableCodec { index_id: IndexId, ordinal: u32, bucket: Option<&Bucket>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -636,7 +638,7 @@ impl TableCodec { table_name: &str, index_id: IndexId, top_n: Option<&ColumnTopN>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -668,7 +670,7 @@ impl TableCodec { } /// Key: `View{BOUND_MIN_TAG}{ViewNameHash}` with encoded view payload. - pub fn with_view_value( + pub fn with_view_value( &mut self, view_name: &str, view: &View, @@ -701,7 +703,7 @@ impl TableCodec { &mut self, table_name: &str, meta: Option<&TableMeta>, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), f: impl FnOnce(&[u8], &[u8]) -> Result, ) -> Result { self.clear_buffers(); @@ -766,14 +768,14 @@ impl TableCodec { index_meta: &IndexMeta, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { index_meta.encode(value, true, reference_tables, arena) } pub fn decode_index_meta( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { IndexMeta::decode::( &mut Cursor::new(bytes), @@ -783,14 +785,26 @@ impl TableCodec { ) } - pub fn decode_index_key( + pub fn decode_index_key( bytes: &[u8], ty: &LogicalType, - mapping: Option>, - ) -> Result { + mapping: &M, + ) -> Result, DatabaseError> { // Hash + TypeTag + Bound Min + Index Id Len + Bound Min let start = TUPLE_KEY_PREFIX_LEN + INDEX_ID_LEN + KEY_BOUND_LEN; - DataValue::memcomparable_decode_mapping(&mut Cursor::new(&bytes[start..]), ty, mapping) + let mut reader = Cursor::new(&bytes[start..]); + let types = match ty { + LogicalType::Tuple(types) => types.as_slice(), + ty => std::slice::from_ref(ty), + }; + let mut values = vec![DataValue::Null; mapping.target_len(types.len())]; + for (index_pos, ty) in types.iter().enumerate() { + let value = DataValue::memcomparable_decode(&mut reader, ty)?; + if let Some(scan_pos) = mapping.scan_index(index_pos) { + values[scan_pos] = value; + } + } + Ok(values) } pub fn decode_index(bytes: &[u8]) -> Result { @@ -801,7 +815,7 @@ impl TableCodec { col: &ColumnCatalog, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { col.encode(value, true, reference_tables, arena) } @@ -809,7 +823,7 @@ impl TableCodec { pub fn decode_column( reader: &mut R, reference_tables: &ReferenceTables, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { // `TableCache` is not theoretically used in `table_collect` because `ColumnCatalog` should not depend on other Column ColumnCatalog::decode::(reader, None, reference_tables, arena) @@ -819,14 +833,14 @@ impl TableCodec { statistics_meta: &StatisticsMetaRoot, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { statistics_meta.encode(value, true, reference_tables, arena) } pub fn decode_statistics_meta( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { StatisticsMetaRoot::decode::( &mut Cursor::new(bytes), @@ -840,14 +854,14 @@ impl TableCodec { sketch_meta: &CountMinSketchMeta, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { sketch_meta.encode(value, true, reference_tables, arena) } pub fn decode_statistics_sketch_meta( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { CountMinSketchMeta::decode::( &mut Cursor::new(bytes), @@ -861,14 +875,14 @@ impl TableCodec { sketch_page: &CountMinSketchPage, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { sketch_page.encode(value, true, reference_tables, arena) } pub fn decode_statistics_sketch_page( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { CountMinSketchPage::decode::( &mut Cursor::new(bytes), @@ -882,7 +896,7 @@ impl TableCodec { bucket: &Bucket, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { bucket.encode(value, true, reference_tables, arena) } @@ -902,7 +916,7 @@ impl TableCodec { pub fn decode_statistics_bucket( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { Bucket::decode::( &mut Cursor::new(bytes), @@ -916,14 +930,14 @@ impl TableCodec { top_n: &ColumnTopN, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { top_n.encode(value, true, reference_tables, arena) } pub fn decode_statistics_top_n( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { ColumnTopN::decode::( &mut Cursor::new(bytes), @@ -948,7 +962,7 @@ impl TableCodec { view: &View, reference_tables: &mut ReferenceTables, bytes: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { bytes.clear(); bytes.resize(4, 0u8); @@ -969,7 +983,7 @@ impl TableCodec { drive: (&T, &TableCache), scala_functions: &ScalaFunctions, table_functions: &TableFunctions, - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { let mut cursor = Cursor::new(bytes); let reference_tables_pos = { @@ -990,14 +1004,14 @@ impl TableCodec { meta: &TableMeta, reference_tables: &mut ReferenceTables, value: &mut Bytes, - arena: &impl MetaArena, + arena: &(impl MetaArena + ?Sized), ) -> Result<(), DatabaseError> { meta.encode(value, true, reference_tables, arena) } pub fn decode_root_table( bytes: &[u8], - arena: &mut impl MetaArena, + arena: &mut (impl MetaArena + ?Sized), ) -> Result { let mut bytes = Cursor::new(bytes); @@ -1586,7 +1600,7 @@ mod tests { let value = Arc::new(value); let index = Index::new( index_id as u32, - &value, + std::slice::from_ref(&value), IndexType::PrimaryKey { is_multiple: false }, ); @@ -1639,7 +1653,7 @@ mod tests { let value = Arc::new(value); let index = Index::new( index_id as u32, - &value, + std::slice::from_ref(&value), IndexType::PrimaryKey { is_multiple: false }, ); diff --git a/src/types/evaluator/binary.rs b/src/types/evaluator/binary.rs index 919b6ee4..a83c7942 100644 --- a/src/types/evaluator/binary.rs +++ b/src/types/evaluator/binary.rs @@ -645,6 +645,8 @@ macro_rules! numeric_binary_evaluator_definition { right: &$crate::types::value::DataValue, ) -> Result<$crate::types::value::DataValue, $crate::errors::DatabaseError> { Ok(match (left, right) { + // Integer remainder by zero would panic; report NULL like division. + ($compute_type(_), $compute_type(v2)) if *v2 == 0 => $crate::types::value::DataValue::Null, ($compute_type(v1), $compute_type(v2)) => $compute_type(*v1 % *v2), ($compute_type(_), $crate::types::value::DataValue::Null) | ($crate::types::value::DataValue::Null, $compute_type(_)) diff --git a/src/types/evaluator/cast.rs b/src/types/evaluator/cast.rs index 12e7bc74..60ec2a43 100644 --- a/src/types/evaluator/cast.rs +++ b/src/types/evaluator/cast.rs @@ -1272,7 +1272,7 @@ mod test { ( CAST_TUPLE * CAST_TYPE_STRIDE + CAST_TUPLE, CastEvaluatorParams::Unit, - DataValue::Tuple(vec![DataValue::Int32(1)], false), + DataValue::Tuple(vec![DataValue::Int32(1)]), ), (u16::MAX, CastEvaluatorParams::Unit, DataValue::Int32(1)), ]; @@ -1847,7 +1847,7 @@ mod test { #[test] fn test_cast_create_dispatches_tuple_and_rejects_unsupported_casts() -> Result<(), DatabaseError> { - let tuple = DataValue::Tuple(vec![DataValue::Int32(1), utf8("2")], false); + let tuple = DataValue::Tuple(vec![DataValue::Int32(1), utf8("2")]); assert_eq!( cast_eval( LogicalType::Tuple(vec![ @@ -1857,7 +1857,7 @@ mod test { LogicalType::Tuple(vec![LogicalType::Bigint, LogicalType::Integer]), &tuple, )?, - DataValue::Tuple(vec![DataValue::Int64(1), DataValue::Int32(2)], false) + DataValue::Tuple(vec![DataValue::Int64(1), DataValue::Int32(2)]) ); assert!(matches!( diff --git a/src/types/evaluator/tuple.rs b/src/types/evaluator/tuple.rs index 5a54220f..16964ef9 100644 --- a/src/types/evaluator/tuple.rs +++ b/src/types/evaluator/tuple.rs @@ -18,10 +18,7 @@ use crate::types::evaluator::DataValue; use std::cmp::Ordering; use std::hint; -fn tuple_cmp( - (v1, v1_is_upper): (&Vec, &bool), - (v2, v2_is_upper): (&Vec, &bool), -) -> Option { +fn tuple_cmp(v1: &[DataValue], v2: &[DataValue]) -> Option { let mut order = Ordering::Equal; let mut v1_iter = v1.iter(); let mut v2_iter = v2.iter(); @@ -29,20 +26,8 @@ fn tuple_cmp( while order == Ordering::Equal { order = match (v1_iter.next(), v2_iter.next()) { (Some(v1), Some(v2)) => v1.partial_cmp(v2)?, - (Some(_), None) => { - if *v2_is_upper { - Ordering::Less - } else { - Ordering::Greater - } - } - (None, Some(_)) => { - if *v1_is_upper { - Ordering::Greater - } else { - Ordering::Less - } - } + (Some(_), None) => Ordering::Greater, + (None, Some(_)) => Ordering::Less, (None, None) => break, } } @@ -53,7 +38,7 @@ pub fn tuple_eq_binary_eval( right: &DataValue, ) -> Result { Ok(match (left, right) { - (DataValue::Tuple(v1, ..), DataValue::Tuple(v2, ..)) => DataValue::Boolean(*v1 == *v2), + (DataValue::Tuple(v1), DataValue::Tuple(v2)) => DataValue::Boolean(*v1 == *v2), (DataValue::Null, DataValue::Boolean(_)) | (DataValue::Boolean(_), DataValue::Null) | (DataValue::Null, DataValue::Null) => DataValue::Null, @@ -65,7 +50,7 @@ pub fn tuple_not_eq_binary_eval( right: &DataValue, ) -> Result { Ok(match (left, right) { - (DataValue::Tuple(v1, ..), DataValue::Tuple(v2, ..)) => DataValue::Boolean(*v1 != *v2), + (DataValue::Tuple(v1), DataValue::Tuple(v2)) => DataValue::Boolean(*v1 != *v2), (DataValue::Null, DataValue::Boolean(_)) | (DataValue::Boolean(_), DataValue::Null) | (DataValue::Null, DataValue::Null) => DataValue::Null, @@ -77,11 +62,9 @@ macro_rules! tuple_order_binary { ($name:ident, $is_order:ident) => { pub fn $name(left: &DataValue, right: &DataValue) -> Result { Ok(match (left, right) { - (DataValue::Tuple(v1, is_upper1), DataValue::Tuple(v2, is_upper2)) => { - tuple_cmp((v1, is_upper1), (v2, is_upper2)) - .map(|order| DataValue::Boolean(order.$is_order())) - .unwrap_or(DataValue::Null) - } + (DataValue::Tuple(v1), DataValue::Tuple(v2)) => tuple_cmp(v1, v2) + .map(|order| DataValue::Boolean(order.$is_order())) + .unwrap_or(DataValue::Null), (DataValue::Null, DataValue::Boolean(_)) | (DataValue::Boolean(_), DataValue::Null) | (DataValue::Null, DataValue::Null) => DataValue::Null, @@ -102,14 +85,14 @@ pub(crate) fn eval_tuple_cast( ) -> Result { match value { DataValue::Null => Ok(DataValue::Null), - DataValue::Tuple(values, is_upper) => { + DataValue::Tuple(values) => { let mut casted = Vec::with_capacity(values.len()); for (value, evaluator) in values.iter().zip(element_evaluators.iter()) { casted.push(evaluator.eval(value)?); } - Ok(DataValue::Tuple(casted, *is_upper)) + Ok(DataValue::Tuple(casted)) } _ => unsafe { hint::unreachable_unchecked() }, } @@ -124,7 +107,7 @@ mod test { use crate::types::LogicalType; fn tuple(values: Vec) -> DataValue { - DataValue::Tuple(values, false) + DataValue::Tuple(values) } #[test] @@ -133,7 +116,7 @@ mod test { let same = tuple(vec![DataValue::Int32(1), DataValue::Int32(2)]); let greater = tuple(vec![DataValue::Int32(1), DataValue::Int32(3)]); let shorter_lower = tuple(vec![DataValue::Int32(1)]); - let shorter_upper = DataValue::Tuple(vec![DataValue::Int32(1)], true); + let shorter_upper = tuple(vec![DataValue::Int32(1)]); let incomparable = tuple(vec![DataValue::Int32(1), DataValue::Boolean(true)]); assert_eq!( @@ -165,7 +148,7 @@ mod test { DataValue::Boolean(true) ); assert_eq!( - tuple_gt_binary_eval(&shorter_upper, &left).unwrap(), + tuple_lt_binary_eval(&shorter_upper, &left).unwrap(), DataValue::Boolean(true) ); assert_eq!( @@ -203,19 +186,16 @@ mod test { assert_eq!( evaluator - .eval(&DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Utf8 { - value: "2".to_string(), - ty: crate::types::value::Utf8Type::Variable(None), - unit: CharLengthUnits::Characters, - }, - ], - false, - )) + .eval(&DataValue::Tuple(vec![ + DataValue::Int32(1), + DataValue::Utf8 { + value: "2".to_string(), + ty: crate::types::value::Utf8Type::Variable(None), + unit: CharLengthUnits::Characters, + }, + ],)) .unwrap(), - DataValue::Tuple(vec![DataValue::Int64(1), DataValue::Int32(2)], false) + DataValue::Tuple(vec![DataValue::Int64(1), DataValue::Int32(2)]) ); assert_eq!(evaluator.eval(&DataValue::Null).unwrap(), DataValue::Null); } diff --git a/src/types/index.rs b/src/types/index.rs index 6e83fee1..bb9bda18 100644 --- a/src/types/index.rs +++ b/src/types/index.rs @@ -18,7 +18,8 @@ use crate::expression::range_detacher::Range; use crate::expression::ScalarExpression; use crate::planner::operator::SortOption; use crate::planner::Explain; -use crate::planner::{ExprRef, PlanArena}; +use crate::planner::ExprRef; +use crate::planner::MetaArena; use crate::types::serialize::TupleValueSerializableImpl; use crate::types::value::DataValue; use crate::types::{ColumnId, LogicalType}; @@ -69,6 +70,18 @@ pub enum IndexLookup { Probe, } +impl IndexLookup { + pub(crate) fn bind_parameters( + &mut self, + params: &[(usize, DataValue)], + ) -> Result<(), DatabaseError> { + if let Self::Static(range) = self { + range.bind_parameters(params)?; + } + Ok(()) + } +} + #[derive(Debug, Clone, Eq, PartialEq, Hash, ReferenceSerialization)] pub struct IndexInfo { pub(crate) meta: IndexMetaRef, @@ -112,44 +125,38 @@ pub struct IndexMeta { } impl IndexMeta { - pub(crate) fn column_exprs( - &self, - table: &TableCatalog, - arena: &PlanArena, - ) -> Result, DatabaseError> { - let mut exprs = Vec::with_capacity(self.column_ids.len()); - - for column_id in self.column_ids.iter() { - if let Some((position, column_ref)) = table + pub(crate) fn column_exprs<'a>( + &'a self, + table: &'a TableCatalog, + ) -> impl Iterator> + 'a { + self.column_ids.iter().copied().map(move |column_id| { + let column_ref = table + .get_column_by_id(&column_id) + .ok_or_else(|| DatabaseError::column_not_found(column_id.to_string()))?; + let position = table .columns() - .copied() - .enumerate() - .find(|(_, column)| arena.column(*column).id() == Some(*column_id)) - { - exprs.push(ScalarExpression::column_expr(column_ref, position)); - } else { - return Err(DatabaseError::column_not_found(column_id.to_string())); - } - } - Ok(exprs) + .position(|column| *column == column_ref) + .ok_or_else(|| DatabaseError::column_not_found(column_id.to_string()))?; + Ok(ScalarExpression::column_expr(column_ref, position)) + }) } } -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Copy)] pub struct Index<'a> { pub id: IndexId, - pub value: &'a DataValue, + pub values: &'a [DataValue], pub ty: IndexType, } impl<'a> Index<'a> { - pub fn new(id: IndexId, value: &'a DataValue, ty: IndexType) -> Self { - Index { id, value, ty } + pub fn new(id: IndexId, values: &'a [DataValue], ty: IndexType) -> Self { + Index { id, values, ty } } } impl Explain for IndexInfo { - fn fmt(&self, arena: &PlanArena<'_>, f: &mut Formatter<'_>) -> fmt::Result { + fn fmt(&self, arena: &dyn MetaArena, f: &mut Formatter<'_>) -> fmt::Result { write!(f, "{} => ", self.meta.explain(arena))?; match &self.lookup { Some(IndexLookup::Static(range)) => write!(f, "{range}")?, @@ -262,9 +269,9 @@ mod tests { assert_eq!(info.explain(&arena).to_string(), "idx_t => 1 Covered"); let value = DataValue::Int32(1); - let index = Index::new(9, &value, IndexType::Unique); + let index = Index::new(9, std::slice::from_ref(&value), IndexType::Unique); assert_eq!(index.id, 9); - assert_eq!(index.value, &value); + assert_eq!(index.values, std::slice::from_ref(&value)); assert_eq!(index.ty, IndexType::Unique); } @@ -408,15 +415,20 @@ mod tests { ty: IndexType::Normal, }; - assert_eq!(meta.column_exprs(&table, &arena)?.len(), 1); + assert_eq!( + meta.column_exprs(&table) + .collect::, _>>()? + .len(), + 1 + ); let missing = IndexMeta { column_ids: vec![u64::MAX], ..meta }; assert!(matches!( - missing.column_exprs(&table, &arena), - Err(DatabaseError::ColumnNotFound { .. }) + missing.column_exprs(&table).next(), + Some(Err(DatabaseError::ColumnNotFound { .. })) )); Ok(()) diff --git a/src/types/tuple.rs b/src/types/tuple.rs index 0e72396c..bd9f07da 100644 --- a/src/types/tuple.rs +++ b/src/types/tuple.rs @@ -15,7 +15,7 @@ use crate::catalog::{ColumnCatalog, ColumnRef}; use crate::errors::DatabaseError; use crate::iter_ext::Itertools; -use crate::planner::PlanArena; +use crate::planner::MetaArena; use crate::types::serialize::{TupleValueSerializable, TupleValueSerializableImpl}; use crate::types::value::DataValue; use std::borrow::Borrow; @@ -28,12 +28,12 @@ pub type Schema = Vec; pub struct SchemaView<'a, 'p> { schema: &'a Schema, - arena: &'a PlanArena<'p>, + arena: &'a (dyn MetaArena + 'p), } pub struct SchemaColumnIter<'a, 'p, 's> { columns: std::slice::Iter<'s, ColumnRef>, - arena: &'a PlanArena<'p>, + arena: &'a (dyn MetaArena + 'p), } impl<'a> Iterator for SchemaColumnIter<'a, '_, '_> { @@ -45,7 +45,7 @@ impl<'a> Iterator for SchemaColumnIter<'a, '_, '_> { } impl<'a, 'p> SchemaView<'a, 'p> { - pub fn new(schema: &'a Schema, arena: &'a PlanArena<'p>) -> Self { + pub fn new(schema: &'a Schema, arena: &'a (dyn MetaArena + 'p)) -> Self { Self { schema, arena } } @@ -265,10 +265,7 @@ impl Tuple { pub fn primary_projection(pk_indices: &[usize], values: &[DataValue]) -> TupleId { if pk_indices.len() > 1 { - DataValue::Tuple( - pk_indices.iter().map(|i| values[*i].clone()).collect_vec(), - false, - ) + DataValue::Tuple(pk_indices.iter().map(|i| values[*i].clone()).collect_vec()) } else { values[pk_indices[0]].clone() } @@ -548,13 +545,10 @@ mod tests { .map(|column| column.datatype().serializable()) .collect_vec(); let mut multi_pk_tuple = tuples[0].clone(); - multi_pk_tuple.pk = Some(DataValue::Tuple( - vec![ - multi_pk_tuple.values[4].clone(), - multi_pk_tuple.values[2].clone(), - ], - false, - )); + multi_pk_tuple.pk = Some(DataValue::Tuple(vec![ + multi_pk_tuple.values[4].clone(), + multi_pk_tuple.values[2].clone(), + ])); let mut tuple_3 = Tuple { pk: multi_pk_tuple.pk.clone(), @@ -696,7 +690,7 @@ mod tests { ); assert_eq!( Tuple::primary_projection(&[0, 1], &tuple.values), - DataValue::Tuple(vec![DataValue::Null, DataValue::Int32(7)], false) + DataValue::Tuple(vec![DataValue::Null, DataValue::Int32(7)]) ); } } diff --git a/src/types/tuple_builder.rs b/src/types/tuple_builder.rs index 82bca29b..f836ec2a 100644 --- a/src/types/tuple_builder.rs +++ b/src/types/tuple_builder.rs @@ -106,10 +106,10 @@ mod tests { let tuple = builder.build_with_row(["7", "kite"]).unwrap(); assert_eq!( tuple.pk, - Some(DataValue::Tuple( - vec![DataValue::Int32(7), DataValue::from("kite".to_string())], - false, - )) + Some(DataValue::Tuple(vec![ + DataValue::Int32(7), + DataValue::from("kite".to_string()) + ])) ); assert_eq!( tuple.values, diff --git a/src/types/value.rs b/src/types/value.rs index 44922b1c..ce836caa 100644 --- a/src/types/value.rs +++ b/src/types/value.rs @@ -15,7 +15,7 @@ use super::LogicalType; use crate::errors::DatabaseError; use crate::iter_ext::Itertools; -use crate::storage::table_codec::{BumpBytes, BOUND_MAX_TAG, NOTNULL_TAG, NULL_TAG}; +use crate::storage::table_codec::{BumpBytes, NOTNULL_TAG, NULL_TAG}; use crate::types::evaluator::cast::{cast_create, to_char, to_varchar}; use crate::types::CharLengthUnits; #[cfg(feature = "time")] @@ -132,6 +132,11 @@ pub enum Utf8Type { #[derive(Clone)] pub enum DataValue { + /// A typed placeholder used only while preparing a plan. + Parameter { + id: usize, + ty: LogicalType, + }, Null, Boolean(bool), Float32(OrderedFloat), @@ -157,8 +162,28 @@ pub enum DataValue { Time64(i64, u64, bool), #[cfg(feature = "decimal")] Decimal(Decimal), - /// (values, is_upper) - Tuple(Vec, bool), + Tuple(Vec), +} + +pub trait IndexKeyMapping { + fn target_len(&self, fields_len: usize) -> usize; + + fn scan_index(&self, index_pos: usize) -> Option; +} + +#[derive(Clone, Copy)] +pub struct OrderedMapping; + +impl IndexKeyMapping for OrderedMapping { + #[inline] + fn target_len(&self, fields_len: usize) -> usize { + fields_len + } + + #[inline] + fn scan_index(&self, index_pos: usize) -> Option { + Some(index_pos) + } } #[derive(Clone, Copy)] @@ -192,42 +217,15 @@ impl<'a> TupleMappingRef<'a> { } } -enum TupleCollector<'a> { - Mapped { - mapping: TupleMappingRef<'a>, - values: Vec, - }, - Ordered(Vec), -} - -impl<'a> TupleCollector<'a> { - fn new(mapping: Option>, tuple_len: usize) -> Self { - if let Some(mapping) = mapping { - TupleCollector::Mapped { - values: vec![DataValue::Null; mapping.target_len()], - mapping, - } - } else { - TupleCollector::Ordered(Vec::with_capacity(tuple_len)) - } - } - - fn push(&mut self, index_pos: usize, value: DataValue) { - match self { - TupleCollector::Mapped { mapping, values } => { - if let Some(target_pos) = mapping.scan_index(index_pos) { - values[target_pos] = value; - } - } - TupleCollector::Ordered(values) => values.push(value), - } +impl IndexKeyMapping for TupleMappingRef<'_> { + #[inline] + fn target_len(&self, _: usize) -> usize { + self.target_len } - fn finish(self) -> Vec { - match self { - TupleCollector::Mapped { values, .. } => values, - TupleCollector::Ordered(values) => values, - } + #[inline] + fn scan_index(&self, index_pos: usize) -> Option { + TupleMappingRef::scan_index(self, index_pos) } } @@ -279,6 +277,17 @@ impl PartialEq for DataValue { } match (self, other) { + ( + Parameter { + id: left, + ty: left_ty, + }, + Parameter { + id: right, + ty: right_ty, + }, + ) => left == right && left_ty == right_ty, + (Parameter { .. }, _) => false, (Boolean(v1), Boolean(v2)) => v1.eq(v2), (Boolean(_), _) => false, (Float32(v1), Float32(v2)) => v1.eq(v2), @@ -317,44 +326,39 @@ impl PartialEq for DataValue { (Decimal(v1), Decimal(v2)) => v1.eq(v2), #[cfg(feature = "decimal")] (Decimal(_), _) => false, - (Tuple(values_1, is_upper_1), Tuple(values_2, is_upper_2)) => { - values_1.eq(values_2) && is_upper_1.eq(is_upper_2) - } + (Tuple(values_1), Tuple(values_2)) => values_1.eq(values_2), (Tuple(..), _) => false, } } } -fn tuple_partial_cmp( - (left, left_is_upper): (&[DataValue], bool), - (right, right_is_upper): (&[DataValue], bool), +pub(crate) fn tuple_partial_cmp( + left: &[DataValue], + right: &[DataValue], + left_is_upper: bool, + right_is_upper: bool, ) -> Option { let mut left_iter = left.iter(); let mut right_iter = right.iter(); loop { - match (left_iter.next(), right_iter.next()) { - (Some(left), Some(right)) => { + match ( + left_iter.next(), + right_iter.next(), + left_is_upper, + right_is_upper, + ) { + (Some(left), Some(right), _, _) => { let ordering = tuple_element_partial_cmp(left, right)?; if ordering != Ordering::Equal { return Some(ordering); } } - (Some(_), None) => { - return Some(if right_is_upper { - Ordering::Less - } else { - Ordering::Greater - }); - } - (None, Some(_)) => { - return Some(if left_is_upper { - Ordering::Greater - } else { - Ordering::Less - }); - } - (None, None) => return Some(Ordering::Equal), + (Some(_), None, _, true) => return Some(Ordering::Less), + (Some(_), None, _, false) => return Some(Ordering::Greater), + (None, Some(_), true, _) => return Some(Ordering::Greater), + (None, Some(_), false, _) => return Some(Ordering::Less), + (None, None, _, _) => return Some(Ordering::Equal), } } } @@ -364,8 +368,8 @@ fn tuple_element_partial_cmp(left: &DataValue, right: &DataValue) -> Option Some(Ordering::Equal), (DataValue::Null, _) => Some(Ordering::Greater), (_, DataValue::Null) => Some(Ordering::Less), - (DataValue::Tuple(left, left_is_upper), DataValue::Tuple(right, right_is_upper)) => { - tuple_partial_cmp((left, *left_is_upper), (right, *right_is_upper)) + (DataValue::Tuple(left), DataValue::Tuple(right)) => { + tuple_partial_cmp(left, right, false, false) } _ => left.partial_cmp(right), } @@ -375,6 +379,7 @@ impl PartialOrd for DataValue { fn partial_cmp(&self, other: &Self) -> Option { use DataValue::*; match (self, other) { + (Parameter { .. }, _) => None, (Boolean(v1), Boolean(v2)) => v1.partial_cmp(v2), (Boolean(_), _) => None, (Float32(v1), Float32(v2)) => v1.partial_cmp(v2), @@ -413,9 +418,7 @@ impl PartialOrd for DataValue { (Decimal(v1), Decimal(v2)) => v1.partial_cmp(v2), #[cfg(feature = "decimal")] (Decimal(_), _) => None, - (Tuple(v1, is_upper1), Tuple(v2, is_upper2)) => { - tuple_partial_cmp((v1, *is_upper1), (v2, *is_upper2)) - } + (Tuple(v1), Tuple(v2)) => tuple_partial_cmp(v1, v2, false, false), (Tuple(..), _) => None, } } @@ -441,6 +444,11 @@ impl Hash for DataValue { fn hash(&self, state: &mut H) { use DataValue::*; match self { + Parameter { id, ty } => { + 19u8.hash(state); + id.hash(state); + ty.hash(state); + } Null => 0u8.hash(state), Boolean(v) => { 1u8.hash(state); @@ -511,17 +519,116 @@ impl Hash for DataValue { 17u8.hash(state); v.hash(state); } - Tuple(values, is_upper) => { + Tuple(values) => { 18u8.hash(state); values.hash(state); - is_upper.hash(state); } } } } impl DataValue { + /// Estimate bound selectivity using the complete index key type. + pub(crate) fn bound_selectivity( + min: std::ops::Bound<&Self>, + max: std::ops::Bound<&Self>, + key_type: &LogicalType, + distinct: f64, + ) -> f64 { + use std::ops::Bound; + + match (min, max, key_type) { + ( + Bound::Included(Self::Tuple(lower)) | Bound::Excluded(Self::Tuple(lower)), + Bound::Included(Self::Tuple(upper)) | Bound::Excluded(Self::Tuple(upper)), + LogicalType::Tuple(fields), + ) if !fields.is_empty() => { + if lower == upper + && lower.len() == fields.len() + && upper.len() == fields.len() + && (matches!(min, Bound::Excluded(_)) || matches!(max, Bound::Excluded(_))) + { + return 0.0; + } + // Approximate field NDV from full-key NDV; no per-field statistics. + let distinct = distinct.powf(1.0 / fields.len() as f64); + let mut divisor = 1.0; + for (i, field) in fields.iter().enumerate() { + let (min, max, is_range) = match (lower.get(i), upper.get(i)) { + (Some(lower), Some(upper)) => ( + Bound::Included(lower), + Bound::Included(upper), + lower != upper, + ), + (Some(lower), None) => (Bound::Included(lower), Bound::Unbounded, true), + (None, Some(upper)) => (Bound::Unbounded, Bound::Included(upper), true), + (None, None) => return 1.0 / divisor, + }; + let selectivity = Self::bound_selectivity(min, max, field, distinct); + if selectivity == 0.0 { + return 0.0; + } + divisor /= selectivity; + // Only the equal prefix and first differing field constrain + // a lexicographic range. Missing fields delimit the prefix. + if is_range { + return 1.0 / divisor; + } + } + 1.0 / divisor + } + (Bound::Included(lower), Bound::Included(upper), _) if lower == upper => 1.0 / distinct, + (Bound::Unbounded, Bound::Unbounded, _) => 1.0, + // SQLite 3.45.3 whereRangeScanEst fallback: 1/4 for one bound, + // 1/64 for two. These are heuristics, not bucket estimates. + (Bound::Unbounded, _, _) | (_, Bound::Unbounded, _) => 1.0 / 4.0, + _ => 1.0 / 64.0, + } + } + + pub(crate) fn parameter_id(placeholder: &str) -> Option { + let digits = placeholder.strip_prefix('$')?; + if digits.is_empty() || digits.starts_with('0') { + return None; + } + digits.parse().ok() + } + + pub(crate) fn bind_parameters( + &mut self, + params: &[(usize, DataValue)], + ) -> Result<(), DatabaseError> { + match self { + DataValue::Parameter { id, ty } => { + let input = params + .iter() + .find_map(|(candidate, value)| (candidate == id).then_some(value)); + if let Some(input) = input { + *self = input.clone().cast(ty)?; + } else { + return Err(DatabaseError::parameter_not_found(format!("${id}"))); + } + } + DataValue::Tuple(values) => { + for value in values { + value.bind_parameters(params)?; + } + } + _ => {} + } + Ok(()) + } + + pub(crate) fn has_parameter(&self) -> bool { + match self { + DataValue::Parameter { .. } => true, + DataValue::Tuple(values) => values.iter().any(Self::has_parameter), + _ => false, + } + } + pub(crate) fn serialized_len_hint(&self) -> usize { match self { + DataValue::Parameter { .. } => 0, DataValue::Null => 0, DataValue::Boolean(_) | DataValue::Int8(_) | DataValue::UInt8(_) => 1, DataValue::Int16(_) | DataValue::UInt16(_) => 2, @@ -545,7 +652,7 @@ impl DataValue { }, #[cfg(feature = "decimal")] DataValue::Decimal(_) => 16, - DataValue::Tuple(values, _) => values.iter().map(DataValue::serialized_len_hint).sum(), + DataValue::Tuple(values) => values.iter().map(DataValue::serialized_len_hint).sum(), } } @@ -774,7 +881,7 @@ impl DataValue { LogicalType::Tuple(types) => { let values = types.iter().map(DataValue::init).collect_vec(); - DataValue::Tuple(values, false) + DataValue::Tuple(values) } } } @@ -782,6 +889,7 @@ impl DataValue { #[inline] pub fn logical_type(&self) -> LogicalType { match self { + DataValue::Parameter { ty, .. } => ty.clone(), DataValue::Null => LogicalType::SqlNull, DataValue::Boolean(_) => LogicalType::Boolean, DataValue::Float32(_) => LogicalType::Float, @@ -810,7 +918,7 @@ impl DataValue { DataValue::Time64(..) => LogicalType::TimeStamp(None, false), #[cfg(feature = "decimal")] DataValue::Decimal(_) => LogicalType::Decimal(None, None), - DataValue::Tuple(values, ..) => { + DataValue::Tuple(values) => { let types = values.iter().map(|v| v.logical_type()).collect_vec(); LogicalType::Tuple(types) } @@ -921,6 +1029,11 @@ impl DataValue { b.push_byte(not_null_tag); match self { + DataValue::Parameter { .. } => { + return Err(DatabaseError::InvalidValue( + "unbound parameter reached storage encoding".to_string(), + )); + } DataValue::Null => (), DataValue::Int8(v) => encode_u!(b, *v as u8 ^ 0x80_u8), DataValue::Int16(v) => encode_u!(b, *v as u16 ^ 0x8000_u16), @@ -961,14 +1074,9 @@ impl DataValue { } #[cfg(feature = "decimal")] DataValue::Decimal(v) => Self::serialize_decimal(*v, b)?, - DataValue::Tuple(values, is_upper) => { - let last = values.len() - 1; - - for (i, v) in values.iter().enumerate() { + DataValue::Tuple(values) => { + for v in values { v.memcomparable_encode(b)?; - if i == last && *is_upper { - b.push_byte(BOUND_MAX_TAG); - } } } } @@ -987,16 +1095,6 @@ impl DataValue { pub fn memcomparable_decode( reader: &mut R, ty: &LogicalType, - ) -> Result { - Self::memcomparable_decode_mapping(reader, ty, None) - } - - #[inline] - pub fn memcomparable_decode_mapping( - reader: &mut R, - ty: &LogicalType, - // for index cover mapping reduce one layer of conversion - tuple_mapping: Option>, ) -> Result { if reader.read_u8()? == NULL_TAG { return Ok(DataValue::Null); @@ -1088,13 +1186,11 @@ impl DataValue { "DECIMAL requires the `decimal` feature".to_string(), )), LogicalType::Tuple(tys) => { - let mut collector = TupleCollector::new(tuple_mapping, tys.len()); - - for (index_pos, ty) in tys.iter().enumerate() { - let value = Self::memcomparable_decode_mapping(reader, ty, None)?; - collector.push(index_pos, value); + let mut values = Vec::with_capacity(tys.len()); + for ty in tys { + values.push(Self::memcomparable_decode(reader, ty)?); } - Ok(DataValue::Tuple(collector.finish(), false)) + Ok(DataValue::Tuple(values)) } } } @@ -1301,6 +1397,9 @@ impl DataValue { } match (self, to) { + (DataValue::Parameter { id, .. }, ty) => { + Ok(DataValue::Parameter { id, ty: ty.clone() }) + } (DataValue::Null, _) => Ok(DataValue::Null), (DataValue::Utf8 { value, .. }, LogicalType::Char(len, unit)) => { to_char(value, *len, *unit) @@ -1622,6 +1721,7 @@ macro_rules! format_float_option { impl fmt::Display for DataValue { fn fmt(&self, f: &mut Formatter) -> fmt::Result { match self { + DataValue::Parameter { id, .. } => write!(f, "${id}")?, DataValue::Boolean(e) => write!(f, "{e}")?, DataValue::Float32(e) => format_float_option!(f, e)?, DataValue::Float64(e) => format_float_option!(f, e)?, @@ -1671,7 +1771,7 @@ impl fmt::Display for DataValue { } #[cfg(feature = "decimal")] DataValue::Decimal(e) => write!(f, "{}", DataValue::decimal_format(e))?, - DataValue::Tuple(values, ..) => { + DataValue::Tuple(values) => { write!(f, "(")?; let len = values.len(); @@ -1691,6 +1791,7 @@ impl fmt::Display for DataValue { impl fmt::Debug for DataValue { fn fmt(&self, f: &mut Formatter) -> fmt::Result { match self { + DataValue::Parameter { ty, .. } => write!(f, "Parameter({self}: {ty})"), DataValue::Boolean(_) => write!(f, "Boolean({self})"), DataValue::Float32(_) => write!(f, "Float32({self})"), DataValue::Float64(_) => write!(f, "Float64({self})"), @@ -1710,13 +1811,7 @@ impl fmt::Debug for DataValue { DataValue::Time64(..) => write!(f, "Time64({self})"), #[cfg(feature = "decimal")] DataValue::Decimal(_) => write!(f, "Decimal({self})"), - DataValue::Tuple(..) => { - write!(f, "Tuple({self}")?; - if matches!(self, DataValue::Tuple(_, true)) { - write!(f, " [is upper]")?; - } - write!(f, ")") - } + DataValue::Tuple(..) => write!(f, "Tuple({self})"), } } } @@ -1726,7 +1821,7 @@ impl fmt::Debug for DataValue { mod test { use crate::errors::DatabaseError; use crate::storage::table_codec::{BumpBytes, NOTNULL_TAG, NULL_TAG}; - use crate::types::value::{DataValue, TupleMappingRef, Utf8Type}; + use crate::types::value::{DataValue, Utf8Type}; use crate::types::CharLengthUnits; use crate::types::LogicalType; use bumpalo::Bump; @@ -1820,8 +1915,7 @@ mod test { assert_eq!(DataValue::Int32(1).serialized_len_hint(), 4); assert_eq!(DataValue::Int64(1).serialized_len_hint(), 8); assert_eq!( - DataValue::Tuple(vec![DataValue::Int8(1), DataValue::Int32(2)], false) - .serialized_len_hint(), + DataValue::Tuple(vec![DataValue::Int8(1), DataValue::Int32(2)]).serialized_len_hint(), 5 ); #[cfg(feature = "decimal")] @@ -1911,7 +2005,7 @@ mod test { assert_eq!( DataValue::init(&LogicalType::Tuple(vec![LogicalType::Integer])), - DataValue::Tuple(vec![DataValue::Int32(0)], false) + DataValue::Tuple(vec![DataValue::Int32(0)]) ); assert_eq!( DataValue::init(&LogicalType::Date).logical_type(), @@ -1983,7 +2077,7 @@ mod test { LogicalType::Decimal(None, None), ), ( - DataValue::Tuple(vec![DataValue::Int32(1)], false), + DataValue::Tuple(vec![DataValue::Int32(1)]), LogicalType::Tuple(vec![LogicalType::Integer]), ), ]; @@ -2181,9 +2275,9 @@ mod test { "Decimal(1.23)", ), ( - DataValue::Tuple(vec![DataValue::Int32(1), utf8("a")], true), + DataValue::Tuple(vec![DataValue::Int32(1), utf8("a")]), "(1, a)", - "Tuple((1, a) [is upper])", + "Tuple((1, a))", ), ( DataValue::Time64(0, 0, false), @@ -2231,14 +2325,9 @@ mod test { #[test] fn test_data_value_nested_tuple_ordering() { assert_eq!( - DataValue::Tuple( - vec![DataValue::Tuple(vec![DataValue::Int32(1)], false)], - false - ) - .partial_cmp(&DataValue::Tuple( - vec![DataValue::Tuple(vec![DataValue::Int32(2)], false)], - false, - )), + DataValue::Tuple(vec![DataValue::Tuple(vec![DataValue::Int32(1)])]).partial_cmp( + &DataValue::Tuple(vec![DataValue::Tuple(vec![DataValue::Int32(2)])],) + ), Some(Ordering::Less) ); } @@ -2400,18 +2489,18 @@ mod test { #[test] fn test_tuple_partial_cmp() { - let tuple_1 = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Int32(2)], false); - let tuple_2 = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Int32(3)], false); - let tuple_with_null = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Null], false); - let lower_prefix = DataValue::Tuple(vec![DataValue::Int32(1)], false); - let upper_prefix = DataValue::Tuple(vec![DataValue::Int32(1)], true); + let tuple_1 = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Int32(2)]); + let tuple_2 = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Int32(3)]); + let tuple_with_null = DataValue::Tuple(vec![DataValue::Int32(1), DataValue::Null]); + let lower_prefix = DataValue::Tuple(vec![DataValue::Int32(1)]); + let upper_prefix = DataValue::Tuple(vec![DataValue::Int32(1)]); assert_eq!(tuple_1.partial_cmp(&tuple_2), Some(Ordering::Less)); assert_eq!(tuple_2.partial_cmp(&tuple_with_null), Some(Ordering::Less)); assert_eq!(lower_prefix.partial_cmp(&tuple_1), Some(Ordering::Less)); - assert_eq!(upper_prefix.partial_cmp(&tuple_1), Some(Ordering::Greater)); + assert_eq!(upper_prefix.partial_cmp(&tuple_1), Some(Ordering::Less)); assert_eq!(tuple_1.partial_cmp(&lower_prefix), Some(Ordering::Greater)); - assert_eq!(tuple_1.partial_cmp(&upper_prefix), Some(Ordering::Less)); + assert_eq!(tuple_1.partial_cmp(&upper_prefix), Some(Ordering::Greater)); assert_eq!(DataValue::Null.partial_cmp(&DataValue::Int32(1)), None); } @@ -2449,8 +2538,8 @@ mod test { DataValue::Decimal(Decimal::new(2, 0)), ), ( - DataValue::Tuple(vec![DataValue::Int32(1)], false), - DataValue::Tuple(vec![DataValue::Int32(2)], false), + DataValue::Tuple(vec![DataValue::Int32(1)]), + DataValue::Tuple(vec![DataValue::Int32(2)]), ), ]; @@ -2491,10 +2580,7 @@ mod test { (DataValue::Time64(1, 0, false), DataValue::Null), #[cfg(feature = "decimal")] (DataValue::Decimal(Decimal::new(1, 0)), DataValue::Null), - ( - DataValue::Tuple(vec![DataValue::Int32(1)], false), - DataValue::Null, - ), + (DataValue::Tuple(vec![DataValue::Int32(1)]), DataValue::Null), ]; for (left, right) in mismatch_cases { @@ -2933,20 +3019,23 @@ mod test { let mut key_tuple_2 = BumpBytes::new_in(&arena); let mut key_tuple_3 = BumpBytes::new_in(&arena); - let v_tuple_1 = DataValue::Tuple( - vec![DataValue::Null, DataValue::Int8(0), DataValue::Int8(1)], - false, - ); + let v_tuple_1 = DataValue::Tuple(vec![ + DataValue::Null, + DataValue::Int8(0), + DataValue::Int8(1), + ]); - let v_tuple_2 = DataValue::Tuple( - vec![DataValue::Int8(0), DataValue::Int8(0), DataValue::Int8(1)], - false, - ); + let v_tuple_2 = DataValue::Tuple(vec![ + DataValue::Int8(0), + DataValue::Int8(0), + DataValue::Int8(1), + ]); - let v_tuple_3 = DataValue::Tuple( - vec![DataValue::Int8(0), DataValue::Int8(0), DataValue::Int8(2)], - false, - ); + let v_tuple_3 = DataValue::Tuple(vec![ + DataValue::Int8(0), + DataValue::Int8(0), + DataValue::Int8(2), + ]); v_tuple_1.memcomparable_encode(&mut key_tuple_1)?; v_tuple_2.memcomparable_encode(&mut key_tuple_2)?; @@ -3001,7 +3090,7 @@ mod test { fn logical_eq(lhs: &DataValue, rhs: &DataValue) -> bool { match (lhs, rhs) { - (Tuple(lv, _), Tuple(rv, _)) => { + (Tuple(lv), Tuple(rv)) => { lv.len() == rv.len() && lv.iter().zip(rv.iter()).all(|(l, r)| logical_eq(l, r)) } _ => lhs == rhs, @@ -3014,14 +3103,11 @@ mod test { let mut key_tuple_2 = BumpBytes::new_in(&arena); let mut key_tuple_3 = BumpBytes::new_in(&arena); - let v_tuple_1 = Tuple( - vec![Null, Int8(0), Int8(1)], - true, // upper bound - ); + let v_tuple_1 = Tuple(vec![Null, Int8(0), Int8(1)]); - let v_tuple_2 = Tuple(vec![Int8(0), Int8(0), Int8(1)], true); + let v_tuple_2 = Tuple(vec![Int8(0), Int8(0), Int8(1)]); - let v_tuple_3 = Tuple(vec![Int8(0), Int8(0), Int8(2)], true); + let v_tuple_3 = Tuple(vec![Int8(0), Int8(0), Int8(2)]); v_tuple_1.memcomparable_encode(&mut key_tuple_1)?; v_tuple_2.memcomparable_encode(&mut key_tuple_2)?; @@ -3047,42 +3133,6 @@ mod test { Ok(()) } - #[test] - fn test_memcomparable_decode_mapping_orders_values() -> Result<(), DatabaseError> { - let arena = Bump::new(); - let mut key_tuple = BumpBytes::new_in(&arena); - - let value = DataValue::Tuple( - vec![ - DataValue::Int32(1), - DataValue::Int32(2), - DataValue::Int32(3), - ], - false, - ); - value.memcomparable_encode(&mut key_tuple)?; - - let ty = LogicalType::Tuple(vec![ - LogicalType::Integer, - LogicalType::Integer, - LogicalType::Integer, - ]); - let index_to_scan = vec![1, usize::MAX, 0]; - let mapping = TupleMappingRef::new(&index_to_scan, 2); - let decoded = DataValue::memcomparable_decode_mapping( - &mut Cursor::new(&key_tuple[..]), - &ty, - Some(mapping), - )?; - - assert_eq!( - decoded, - DataValue::Tuple(vec![DataValue::Int32(3), DataValue::Int32(1)], false) - ); - - Ok(()) - } - #[test] fn test_mem_comparable_utf8() -> Result<(), DatabaseError> { let arena = Bump::new(); diff --git a/src/wasm.rs b/src/wasm.rs index 9380be87..9d786f24 100644 --- a/src/wasm.rs +++ b/src/wasm.rs @@ -33,6 +33,9 @@ fn set_prop(object: &Object, key: &str, value: JsValue) -> Result<(), JsValue> { fn data_value_to_js(value: &DataValue) -> Result { match value { + DataValue::Parameter { id, .. } => Err(to_js_err(format!( + "unbound parameter ${id} reached WASM output" + ))), DataValue::Null => Ok(JsValue::NULL), DataValue::Boolean(value) => Ok(JsValue::from_bool(*value)), DataValue::Float32(value) => Ok(JsValue::from_f64(value.0 as f64)), @@ -69,10 +72,9 @@ fn data_value_to_js(value: &DataValue) -> Result { } #[cfg(feature = "decimal")] DataValue::Decimal(value) => Ok(JsValue::from_str(&value.to_string())), - DataValue::Tuple(values, is_upper) => { + DataValue::Tuple(values) => { let object = Object::new(); set_prop(&object, "values", data_values_to_js(values)?)?; - set_prop(&object, "isUpper", JsValue::from_bool(*is_upper))?; Ok(object.into()) } } diff --git a/tests/slt/group_by.slt b/tests/slt/group_by.slt index d7fdc691..29f6c9e9 100644 --- a/tests/slt/group_by.slt +++ b/tests/slt/group_by.slt @@ -57,5 +57,11 @@ select v2, v2 + 1, sum(v1) from t group by v2 + 1, v2 order by v2 # 6 # 7 +query II +select id + 1, max(v1 + 2) from t where v1 > 3 group by id + 1 order by id + 1; +---- +4 6 +5 7 + statement ok drop table t diff --git a/tests/slt/values.slt b/tests/slt/values.slt index b0c8febc..990eb21d 100644 --- a/tests/slt/values.slt +++ b/tests/slt/values.slt @@ -69,3 +69,23 @@ SELECT t.y FROM (VALUES (10, 20), (30, 40)) AS t(x, y); ---- 20 40 + +statement ok +ALTER TABLE t ADD COLUMN v INT NOT NULL DEFAULT 7; + +statement ok +INSERT INTO t VALUES (4, DEFAULT), (4 + 1, 21), (10, DEFAULT), (10 + 1, 30); + +query II +SELECT * FROM t ORDER BY x; +---- +1 7 +2 7 +3 7 +4 7 +5 21 +10 7 +11 30 + +statement ok +DROP TABLE t; diff --git a/tests/slt/where_by_index.slt b/tests/slt/where_by_index.slt index 902a2ae4..b87056cf 100644 --- a/tests/slt/where_by_index.slt +++ b/tests/slt/where_by_index.slt @@ -304,6 +304,47 @@ select c2, c3 from t_cover where c1 = 2; statement ok drop table t_cover; +statement ok +create table tuple_bounds(w int, k int, primary key(w,k)); + +statement ok +insert into tuple_bounds values(1,1),(1,2),(1,3),(2,1),(2,2),(2,4); + +# Populate other prefixes for ANALYZE without changing the queried rows. +statement ok +insert into tuple_bounds select a.x + 3, b.x +from (values (0),(1),(2),(3),(4),(5),(6),(7),(8),(9)) as a(x) +cross join (values (0),(1),(2),(3),(4),(5),(6),(7),(8),(9)) as b(x); + +statement ok +analyze table tuple_bounds; + +query T +explain select k from tuple_bounds where w = 2 order by k; +---- +Projection [tuple_bounds.k] [Project => (Sort Option: Follow)] TableScan tuple_bounds -> [tuple_bounds.w, tuple_bounds.k] [IndexScan By pk_index => [(2), (2)] Covered => (Sort Option: OrderBy: (tuple_bounds.w Asc Nulls Last, tuple_bounds.k Asc Nulls Last) ignore_prefix_len: 1)] + +query I +select k from tuple_bounds where w = 2 order by k; +---- +1 +2 +4 + +query T +explain select k from tuple_bounds where w = 1 and k >= 1 and k < 3 order by k; +---- +Projection [tuple_bounds.k] [Project => (Sort Option: Follow)] TableScan tuple_bounds -> [tuple_bounds.w, tuple_bounds.k] [IndexScan By pk_index => [(1, 1), (1, 3)) Covered => (Sort Option: OrderBy: (tuple_bounds.w Asc Nulls Last, tuple_bounds.k Asc Nulls Last) ignore_prefix_len: 1)] + +query I +select k from tuple_bounds where w = 1 and k >= 1 and k < 3 order by k; +---- +1 +2 + +statement ok +drop table tuple_bounds; + # Composite index range bounds # Composite equality-prefix bounds must include the complete matching prefix # and distinguish < / <= / > / >= when trailing index columns are present. diff --git a/tests/slt/where_by_index_explain.slt b/tests/slt/where_by_index_explain.slt index 1a629d84..f37d43c0 100644 --- a/tests/slt/where_by_index_explain.slt +++ b/tests/slt/where_by_index_explain.slt @@ -19,6 +19,17 @@ create index p_index on t1 (c1, c2); statement ok analyze table t1; +# Alias substitution after column pruning must retain index access on both sides. +query T +explain select a.c2, b.c2 from t1 a join t1 b on a.id + 3 = b.id where a.id = 3 and b.id = 6; +---- +Projection [a.c2, b.c2] [Project => (Sort Option: Follow)] Inner Join Where ((a.id + 3) = b.id) [NestLoopJoin => (Sort Option: None)] TableScan t1 -> [t1.id, t1.c2] [IndexScan By pk_index => 3 => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] Filter ((t1.id) as (b.id) = 6), Is Having: false [Filter => (Sort Option: Follow)] TableScan t1 -> [t1.id, t1.c2] [IndexScan By pk_index => 6 => (Sort Option: OrderBy: (t1.id Asc Nulls Last) ignore_prefix_len: 0)] + +query II +select a.c2, b.c2 from t1 a join t1 b on a.id + 3 = b.id where a.id = 3 and b.id = 6; +---- +5 8 + query T explain select * from t1 limit 10; ---- diff --git a/tests/sqllogictest/src/lib.rs b/tests/sqllogictest/src/lib.rs index f379fb35..69ca8eda 100644 --- a/tests/sqllogictest/src/lib.rs +++ b/tests/sqllogictest/src/lib.rs @@ -32,37 +32,36 @@ impl DB for SQLBase { println!("|— Input SQL: {}", sql); let mut statements = prepare_all(sql)?.into_iter().peekable(); - while let Some(statement) = statements.next() { + let output = loop { + let Some(statement) = statements.next() else { + break DBOutput::StatementComplete(0); + }; let is_last = statements.peek().is_none(); match command_type(&statement)? { CommandType::DDL => { self.db.ddl(statement.to_string())?; if is_last { - println!(" |— time spent: {:?}", start.elapsed()); - return Ok(DBOutput::StatementComplete(0)); + break DBOutput::StatementComplete(0); } } CommandType::Analyze => { execute_analyze_statement(&mut self.db, &statement)?; if is_last { - println!(" |— time spent: {:?}", start.elapsed()); - return Ok(DBOutput::StatementComplete(0)); + break DBOutput::StatementComplete(0); } } _ => { - let iter = (&self.db).execute(statement, &[])?; + let iter = self.db.run(statement.to_string())?; if is_last { - let output = collect_output(iter)?; - println!(" |— time spent: {:?}", start.elapsed()); - return Ok(output); + break collect_output(iter)?; } iter.done()?; } } - } + }; println!(" |— time spent: {:?}", start.elapsed()); - Ok(DBOutput::StatementComplete(0)) + Ok(output) } } diff --git a/tpcc/README.md b/tpcc/README.md index 1181f0fc..e0ccfb96 100644 --- a/tpcc/README.md +++ b/tpcc/README.md @@ -46,12 +46,12 @@ Local 720-second comparison on the machine above: | Backend | TpmC | New-Order p90 (µs) | Payment p90 (µs) | Order-Status p90 (µs) | Delivery p90 (µs) | Stock-Level p90 (µs) | | --- | ---: | ---: | ---: | ---: | ---: | ---: | -| KiteSQL LMDB | 70315 | 561 | 196 | 208 | 1089 | 1693 | -| KiteSQL RocksDB | 35556 | 747 | 394 | 365 | 10175 | 2125 | -| SQLite balanced | 54797 | 303 | 72 | 52 | 363 | 508 | -| SQLite practical | 44847 | 527 | 146 | 61 | 1248 | 486 | +| KiteSQL LMDB | 128127 | 382 | 100 | 151 | 539 | 490 | +| KiteSQL RocksDB | 42930 | 539 | 299 | 278 | 11743 | 963 | +| SQLite balanced | 55455 | 299 | 71 | 50 | 365 | 519 | +| SQLite practical | 42581 | 365 | 74 | 44 | 485 | 378 | -- Run dates: `2026-09-06–2026-09-07`; results: `2026-09-06_19-04-32`. Latency is measured in microseconds and includes commit. +- Run dates: `2026-09-26`; results: `2026-09-26_13-45-51`. Latency is measured in microseconds and includes commit. - All rows use `--num-ware 1`, `--max-retry 5`, and TPCC's default 720-second measure time. - SQLite rows use the `balanced` and `practical` profiles respectively. diff --git a/tpcc/src/backend/dual.rs b/tpcc/src/backend/dual.rs index a197233f..782dcc1b 100644 --- a/tpcc/src/backend/dual.rs +++ b/tpcc/src/backend/dual.rs @@ -28,7 +28,7 @@ pub struct DualBackend { } pub struct DualPreparedStatement<'a> { - kitesql: KiteSqlPreparedStatement, + kitesql: KiteSqlPreparedStatement<'a>, sqlite: SqlitePreparedStatement<'a>, spec: StatementSpec, } diff --git a/tpcc/src/backend/kitesql_lmdb.rs b/tpcc/src/backend/kitesql_lmdb.rs index bed4101a..bb433500 100644 --- a/tpcc/src/backend/kitesql_lmdb.rs +++ b/tpcc/src/backend/kitesql_lmdb.rs @@ -18,7 +18,7 @@ use super::{ StatementSpec, }; use crate::TpccError; -use kite_sql::db::{prepare, DBTransaction, DataBaseBuilder, Database}; +use kite_sql::db::{DBTransaction, DataBaseBuilder, Database}; use kite_sql::storage::lmdb::{LmdbStorage, LmdbTransaction as KiteSqlLmdbTransaction}; use kite_sql::types::tuple::Tuple; @@ -38,14 +38,13 @@ impl KiteSqlLmdbBackend { fn prepare_spec_groups( &self, specs: &[Vec], - ) -> Result>, TpccError> { + ) -> Result>>, TpccError> { let mut groups = Vec::with_capacity(specs.len()); for group in specs { let mut prepared = Vec::with_capacity(group.len()); for spec in group { - let statement = prepare(spec.sql)?; prepared.push(KiteSqlPreparedStatement { - statement, + plan: self.database.prepare(spec.sql, &spec.parameters)?, spec: spec.clone(), }); } @@ -63,7 +62,7 @@ impl KiteSqlLmdbBackend { impl BackendControl for KiteSqlLmdbBackend { type PreparedStatement<'a> - = KiteSqlPreparedStatement + = KiteSqlPreparedStatement<'a> where Self: 'a; @@ -97,17 +96,17 @@ pub struct KiteSqlLmdbTransactionWrapper<'a> { impl<'a> KiteSqlLmdbTransactionWrapper<'a> { pub(crate) fn execute_raw<'b>( &'b mut self, - statement: &mut KiteSqlPreparedStatement, + statement: &'b mut KiteSqlPreparedStatement<'a>, params: &[DbParam], ) -> Result>, TpccError> { Ok(KiteSqlTxnResult::new( - self.inner.execute(&statement.statement, params)?, + self.inner.execute(&statement.plan, params)?, )) } } impl<'a> BackendTransaction for KiteSqlLmdbTransactionWrapper<'a> { - type PreparedStatement = KiteSqlPreparedStatement; + type PreparedStatement = KiteSqlPreparedStatement<'a>; fn execute_drain( &mut self, diff --git a/tpcc/src/backend/kitesql_rocksdb.rs b/tpcc/src/backend/kitesql_rocksdb.rs index 9a9cd82e..50c1d575 100644 --- a/tpcc/src/backend/kitesql_rocksdb.rs +++ b/tpcc/src/backend/kitesql_rocksdb.rs @@ -19,7 +19,7 @@ use super::{ use crate::TpccError; use kite_sql::binder::{command_type, CommandType}; use kite_sql::db::{ - prepare, prepare_all, DBTransaction, DataBaseBuilder, Database, Statement, TransactionIter, + prepare_all, DBTransaction, DataBaseBuilder, Database, Statement, TransactionIter, }; use kite_sql::storage::rocksdb::{OptimisticRocksStorage, RocksStorage}; use kite_sql::storage::{Storage, Transaction}; @@ -57,14 +57,13 @@ impl KiteSqlRocksBackend { fn prepare_spec_groups( &self, specs: &[Vec], - ) -> Result>, TpccError> { + ) -> Result>>, TpccError> { let mut groups = Vec::with_capacity(specs.len()); for group in specs { let mut prepared = Vec::with_capacity(group.len()); for spec in group { - let statement = prepare(spec.sql)?; prepared.push(KiteSqlPreparedStatement { - statement, + plan: self.database.prepare(spec.sql, &spec.parameters)?, spec: spec.clone(), }); } @@ -82,7 +81,7 @@ impl KiteSqlRocksBackend { impl BackendControl for KiteSqlRocksBackend { type PreparedStatement<'a> - = KiteSqlPreparedStatement + = KiteSqlPreparedStatement<'a> where Self: 'a; @@ -156,9 +155,7 @@ pub(crate) fn execute_kitesql_batch( } let mut transaction = database.new_transaction()?; - for statement in statements { - transaction.execute(statement, &[])?.done()?; - } + transaction.run(sql)?.done()?; transaction.commit()?; Ok(()) } @@ -170,17 +167,17 @@ pub struct KiteSqlRocksTransaction<'a, S: Storage> { impl<'a, S: Storage> KiteSqlRocksTransaction<'a, S> { pub(crate) fn execute_raw<'b>( &'b mut self, - statement: &mut KiteSqlPreparedStatement, + statement: &'b mut KiteSqlPreparedStatement<'a>, params: &[DbParam], ) -> Result>, TpccError> { Ok(KiteSqlTxnResult::new( - self.inner.execute(&statement.statement, params)?, + self.inner.execute(&statement.plan, params)?, )) } } impl<'a, S: Storage> BackendTransaction for KiteSqlRocksTransaction<'a, S> { - type PreparedStatement = KiteSqlPreparedStatement; + type PreparedStatement = KiteSqlPreparedStatement<'a>; fn execute_drain( &mut self, diff --git a/tpcc/src/backend/mod.rs b/tpcc/src/backend/mod.rs index 35df2c5d..09688b0a 100644 --- a/tpcc/src/backend/mod.rs +++ b/tpcc/src/backend/mod.rs @@ -18,11 +18,12 @@ pub mod kitesql_rocksdb; pub mod sqlite; use crate::TpccError; -use kite_sql::db::Statement; +use kite_sql::db::PreparedPlan; use kite_sql::types::tuple::Tuple; use kite_sql::types::value::DataValue; +use kite_sql::types::LogicalType; -pub type DbParam = (&'static str, DataValue); +pub type DbParam = (usize, DataValue); pub trait SimpleExecutor { fn execute_batch(&mut self, sql: &str) -> Result<(), TpccError>; @@ -82,27 +83,15 @@ pub trait BackendTransaction { fn commit(self) -> Result<(), TpccError>; } -#[derive(Clone, Copy)] -pub enum ColumnType { - Int8, - Int16, - Int32, - Int64, - Decimal, - Utf8, - DateTime, - NullableDateTime, -} - #[derive(Clone)] pub struct StatementSpec { pub sql: &'static str, - pub result_types: &'static [ColumnType], + pub result_types: Vec, + pub parameters: Vec<(usize, LogicalType)>, } -#[derive(Clone)] -pub struct KiteSqlPreparedStatement { - pub statement: Statement, +pub struct KiteSqlPreparedStatement<'a> { + pub plan: PreparedPlan<'a>, pub spec: StatementSpec, } @@ -110,7 +99,7 @@ pub trait PreparedStatement { fn spec(&self) -> &StatementSpec; } -impl PreparedStatement for KiteSqlPreparedStatement { +impl PreparedStatement for KiteSqlPreparedStatement<'_> { fn spec(&self) -> &StatementSpec { &self.spec } diff --git a/tpcc/src/backend/sqlite.rs b/tpcc/src/backend/sqlite.rs index c4c5b8dc..5e4dd25a 100644 --- a/tpcc/src/backend/sqlite.rs +++ b/tpcc/src/backend/sqlite.rs @@ -13,15 +13,14 @@ // limitations under the License. use super::{ - BackendControl, BackendTransaction, ColumnType, DbParam, PreparedStatement, SimpleExecutor, - StatementSpec, + BackendControl, BackendTransaction, DbParam, PreparedStatement, SimpleExecutor, StatementSpec, }; use crate::TpccError; use chrono::{NaiveDateTime, TimeZone, Utc}; use clap::ValueEnum; use kite_sql::types::tuple::Tuple; use kite_sql::types::value::{DataValue, Utf8Type}; -use kite_sql::types::CharLengthUnits; +use kite_sql::types::{CharLengthUnits, LogicalType}; use rust_decimal::Decimal; use sqlite::{Connection, State, Statement as SqliteStatement, Value}; @@ -133,7 +132,7 @@ impl<'a> SqliteTransaction<'a> { ) -> Result, TpccError> { statement.statement.reset()?; bind_params(&mut statement.statement, params)?; - SqliteResult::new(&mut statement.statement, statement.spec.result_types) + SqliteResult::new(&mut statement.statement, &statement.spec.result_types) } } @@ -209,20 +208,15 @@ impl<'a> BackendTransaction for SqliteTransaction<'a> { } fn bind_params(statement: &mut SqliteStatement<'_>, params: &[DbParam]) -> Result<(), TpccError> { - for (key, value) in params { - let sqlite_value = convert_value(value)?; - if let Some(index) = key.strip_prefix('?') { - let idx: usize = index.parse().map_err(|_| TpccError::InvalidParameter)?; - statement.bind((idx, sqlite_value.clone()))?; - } else { - statement.bind((key.as_ref(), sqlite_value.clone()))?; - } + for (id, value) in params { + statement.bind((*id, convert_value(value)?))?; } Ok(()) } fn convert_value(value: &DataValue) -> Result { Ok(match value { + DataValue::Parameter { .. } => return Err(TpccError::InvalidParameter), DataValue::Null => Value::Null, DataValue::Boolean(v) => Value::Integer(*v as i64), DataValue::Float32(v) => Value::Float(v.0 as f64), @@ -247,7 +241,7 @@ fn convert_value(value: &DataValue) -> Result { DataValue::Time32(_, _) => Value::Null, DataValue::Time64(value, precision, _) => Value::String(format_time64(*value, *precision)?), DataValue::Decimal(v) => Value::String(v.to_string()), - DataValue::Tuple(_, _) => Value::Null, + DataValue::Tuple(_) => Value::Null, }) } @@ -321,13 +315,13 @@ fn normalize_sqlite_sql(sql: &str) -> Option { pub struct SqliteResult<'stmt, 'conn> { statement: &'stmt mut SqliteStatement<'conn>, - column_types: &'static [ColumnType], + column_types: &'stmt [LogicalType], } impl<'stmt, 'conn> SqliteResult<'stmt, 'conn> { fn new( statement: &'stmt mut SqliteStatement<'conn>, - column_types: &'static [ColumnType], + column_types: &'stmt [LogicalType], ) -> Result { Ok(Self { statement, @@ -356,32 +350,34 @@ impl Iterator for SqliteResult<'_, '_> { fn convert_statement_row( statement: &SqliteStatement<'_>, - types: &[ColumnType], + types: &[LogicalType], ) -> Result { let mut values = Vec::with_capacity(types.len()); for (idx, column_type) in types.iter().enumerate() { let value = match column_type { - ColumnType::Int8 => DataValue::Int8(statement.read::(idx)? as i8), - ColumnType::Int16 => DataValue::Int16(statement.read::(idx)? as i16), - ColumnType::Int32 => DataValue::Int32(statement.read::(idx)? as i32), - ColumnType::Int64 => DataValue::Int64(statement.read::(idx)?), - ColumnType::Decimal => DataValue::Decimal(read_decimal(statement, idx)?), - ColumnType::Utf8 => DataValue::Utf8 { + LogicalType::Tinyint => DataValue::Int8(statement.read::(idx)? as i8), + LogicalType::Smallint => DataValue::Int16(statement.read::(idx)? as i16), + LogicalType::Integer => DataValue::Int32(statement.read::(idx)? as i32), + LogicalType::Bigint => DataValue::Int64(statement.read::(idx)?), + LogicalType::Decimal(..) => DataValue::Decimal(read_decimal(statement, idx)?), + LogicalType::Varchar(..) | LogicalType::Char(..) => DataValue::Utf8 { value: statement.read::(idx)?, ty: Utf8Type::Variable(None), unit: CharLengthUnits::Characters, }, - ColumnType::DateTime => { - let text: String = statement.read(idx)?; - parse_datetime(&text)? - } - ColumnType::NullableDateTime => { + LogicalType::DateTime => { let text: Option = statement.read(idx)?; match text { Some(value) => parse_datetime(&value)?, None => DataValue::Null, } } + _ => { + return Err(kite_sql::errors::DatabaseError::InvalidValue(format!( + "unsupported TPCC result type: {column_type}" + )) + .into()); + } }; values.push(value); } diff --git a/tpcc/src/delivery.rs b/tpcc/src/delivery.rs index 927d6d55..23fd022c 100644 --- a/tpcc/src/delivery.rs +++ b/tpcc/src/delivery.rs @@ -51,8 +51,8 @@ impl TpccTransaction for Delivery { tx.with_query_one( &mut statements[0], &[ - ("$1", DataValue::Int8(d_id as i8)), - ("$2", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int8(d_id as i8)), + (2, DataValue::Int16(args.w_id as i16)), ], &mut |tuple| { no_o_id = tuple.values[0].i32().unwrap(); @@ -67,9 +67,9 @@ impl TpccTransaction for Delivery { tx.execute_drain( &mut statements[1], &[ - ("$1", DataValue::Int32(no_o_id)), - ("$2", DataValue::Int8(d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int32(no_o_id)), + (2, DataValue::Int8(d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), ], )?; // "SELECT o_c_id FROM orders WHERE o_id = ? AND o_d_id = ? AND o_w_id = ?" @@ -77,9 +77,9 @@ impl TpccTransaction for Delivery { tx.with_query_one( &mut statements[2], &[ - ("$1", DataValue::Int32(no_o_id)), - ("$2", DataValue::Int8(d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int32(no_o_id)), + (2, DataValue::Int8(d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), ], &mut |tuple| { c_id = tuple.values[0].i32().unwrap(); @@ -90,20 +90,20 @@ impl TpccTransaction for Delivery { tx.execute_drain( &mut statements[3], &[ - ("$1", DataValue::Int8(args.o_carrier_id as i8)), - ("$2", DataValue::Int32(no_o_id)), - ("$3", DataValue::Int8(d_id as i8)), - ("$4", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int8(args.o_carrier_id as i8)), + (2, DataValue::Int32(no_o_id)), + (3, DataValue::Int8(d_id as i8)), + (4, DataValue::Int16(args.w_id as i16)), ], )?; // "UPDATE order_line SET ol_delivery_d = ? WHERE ol_o_id = ? AND ol_d_id = ? AND ol_w_id = ?" tx.execute_drain( &mut statements[4], &[ - ("$1", DataValue::from(&now)), - ("$2", DataValue::Int32(no_o_id)), - ("$3", DataValue::Int8(d_id as i8)), - ("$4", DataValue::Int16(args.w_id as i16)), + (1, DataValue::from(&now)), + (2, DataValue::Int32(no_o_id)), + (3, DataValue::Int8(d_id as i8)), + (4, DataValue::Int16(args.w_id as i16)), ], )?; // "SELECT SUM(ol_amount) FROM order_line WHERE ol_o_id = ? AND ol_d_id = ? AND ol_w_id = ?" @@ -111,9 +111,9 @@ impl TpccTransaction for Delivery { tx.with_query_one( &mut statements[5], &[ - ("$1", DataValue::Int32(no_o_id)), - ("$2", DataValue::Int8(d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int32(no_o_id)), + (2, DataValue::Int8(d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), ], &mut |tuple| { ol_total = tuple.values[0].decimal().unwrap(); @@ -124,10 +124,10 @@ impl TpccTransaction for Delivery { tx.execute_drain( &mut statements[6], &[ - ("$1", DataValue::Decimal(ol_total)), - ("$2", DataValue::Int32(c_id)), - ("$3", DataValue::Int8(d_id as i8)), - ("$4", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Decimal(ol_total)), + (2, DataValue::Int32(c_id)), + (3, DataValue::Int8(d_id as i8)), + (4, DataValue::Int16(args.w_id as i16)), ], )?; } diff --git a/tpcc/src/main.rs b/tpcc/src/main.rs index 2cf3b73a..a8385514 100644 --- a/tpcc/src/main.rs +++ b/tpcc/src/main.rs @@ -16,7 +16,7 @@ use crate::backend::dual::DualBackend; use crate::backend::kitesql_lmdb::KiteSqlLmdbBackend; use crate::backend::kitesql_rocksdb::{KiteSqlOptimisticRocksDbBackend, KiteSqlRocksDbBackend}; use crate::backend::sqlite::{SqliteBackend, SqliteProfile}; -use crate::backend::{BackendControl, BackendTransaction, ColumnType, StatementSpec}; +use crate::backend::{BackendControl, BackendTransaction, StatementSpec}; use crate::delivery::DeliveryTest; use crate::load::Load; use crate::new_ord::NewOrdTest; @@ -28,6 +28,7 @@ use crate::utils::SeqGen; use clap::{Parser, ValueEnum}; use indicatif::{ProgressBar, ProgressStyle}; use kite_sql::errors::DatabaseError; +use kite_sql::types::LogicalType; #[cfg(all(unix, feature = "pprof"))] use pprof::ProfilerGuard; use rand::prelude::ThreadRng; @@ -364,181 +365,209 @@ impl PprofSession { } fn statement_specs() -> Vec> { + use kite_sql::types::LogicalType::*; vec![ vec![ stmt( "SELECT c.c_discount, c.c_last, c.c_credit, w.w_tax FROM customer AS c JOIN warehouse AS w ON c.c_w_id = w_id AND w.w_id = $1 AND c.c_w_id = $2 AND c.c_d_id = $3 AND c.c_id = $4", - &[ColumnType::Decimal, ColumnType::Utf8, ColumnType::Utf8, ColumnType::Decimal], + vec![Smallint, Smallint, Tinyint, Bigint], + vec![Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters), Varchar(None, kite_sql::types::CharLengthUnits::Characters), Decimal(None, None)], ), stmt( "SELECT c_discount, c_last, c_credit FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_id = $3", - &[ColumnType::Decimal, ColumnType::Utf8, ColumnType::Utf8], + vec![Smallint, Tinyint, Integer], + vec![Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters), Varchar(None, kite_sql::types::CharLengthUnits::Characters)], ), stmt( "SELECT w_tax FROM warehouse WHERE w_id = $1", - &[ColumnType::Decimal], + vec![Smallint], + vec![Decimal(None, None)], ), stmt( "SELECT d_next_o_id, d_tax FROM district WHERE d_id = $1 AND d_w_id = $2", - &[ColumnType::Int32, ColumnType::Decimal], + vec![Tinyint, Smallint], + vec![Integer, Decimal(None, None)], ), stmt( "UPDATE district SET d_next_o_id = $1 + 1 WHERE d_id = $2 AND d_w_id = $3", - &[], + vec![Integer, Tinyint, Smallint], + vec![], ), stmt( "INSERT INTO orders (o_id, o_d_id, o_w_id, o_c_id, o_entry_d, o_ol_cnt, o_all_local) VALUES($1, $2, $3, $4, $5, $6, $7)", - &[], + vec![Integer, Tinyint, Smallint, Integer, DateTime, Tinyint, Tinyint], + vec![], ), stmt( "INSERT INTO new_orders (no_o_id, no_d_id, no_w_id) VALUES ($1,$2,$3)", - &[], + vec![Integer, Tinyint, Smallint], + vec![], ), stmt( "SELECT i_price, i_name, i_data FROM item WHERE i_id = $1", - &[ColumnType::Decimal, ColumnType::Utf8, ColumnType::Utf8], + vec![Integer], + vec![Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters), Varchar(None, kite_sql::types::CharLengthUnits::Characters)], ), stmt( "SELECT s_quantity, s_data, s_dist_01, s_dist_02, s_dist_03, s_dist_04, s_dist_05, s_dist_06, s_dist_07, s_dist_08, s_dist_09, s_dist_10 FROM stock WHERE s_i_id = $1 AND s_w_id = $2", - &[ - ColumnType::Int16, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, + vec![Integer, Smallint], + vec![ + Smallint, + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), ], ), stmt( "UPDATE stock SET s_quantity = $1 WHERE s_i_id = $2 AND s_w_id = $3", - &[], + vec![Smallint, Integer, Smallint], + vec![], ), stmt( "INSERT INTO order_line (ol_o_id, ol_d_id, ol_w_id, ol_number, ol_i_id, ol_supply_w_id, ol_quantity, ol_amount, ol_dist_info) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)", - &[], + vec![Integer, Tinyint, Smallint, Tinyint, Integer, Smallint, Tinyint, Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![], ), ], vec![ stmt( "UPDATE warehouse SET w_ytd = w_ytd + $1 WHERE w_id = $2", - &[], + vec![Decimal(None, None), Smallint], + vec![], ), stmt( "SELECT w_street_1, w_street_2, w_city, w_state, w_zip, w_name FROM warehouse WHERE w_id = $1", - &[ - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, + vec![Smallint], + vec![ + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), ], ), stmt( "UPDATE district SET d_ytd = d_ytd + $1 WHERE d_w_id = $2 AND d_id = $3", - &[], + vec![Decimal(None, None), Smallint, Tinyint], + vec![], ), stmt( "SELECT d_street_1, d_street_2, d_city, d_state, d_zip, d_name FROM district WHERE d_w_id = $1 AND d_id = $2", - &[ - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, + vec![Smallint, Tinyint], + vec![ + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), ], ), stmt( "SELECT count(c_id) FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_last = $3", - &[ColumnType::Int32], + vec![Smallint, Tinyint, Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![Integer], ), stmt( "SELECT c_id FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_last = $3 ORDER BY c_first", - &[ColumnType::Int32], + vec![Smallint, Tinyint, Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![Integer], ), stmt( "SELECT c_first, c_middle, c_last, c_street_1, c_street_2, c_city, c_state, c_zip, c_phone, c_credit, c_credit_lim, c_discount, c_balance, c_since FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_id = $3", - &[ - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Int64, - ColumnType::Decimal, - ColumnType::Decimal, - ColumnType::DateTime, + vec![Smallint, Tinyint, Integer], + vec![ + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Bigint, + Decimal(None, None), + Decimal(None, None), + DateTime, ], ), stmt( "SELECT c_data FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_id = $3", - &[ColumnType::Utf8], + vec![Smallint, Tinyint, Integer], + vec![Varchar(None, kite_sql::types::CharLengthUnits::Characters)], ), stmt( "UPDATE customer SET c_balance = $1, c_data = $2 WHERE c_w_id = $3 AND c_d_id = $4 AND c_id = $5", - &[], + vec![Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters), Smallint, Tinyint, Integer], + vec![], ), stmt( "UPDATE customer SET c_balance = $1 WHERE c_w_id = $2 AND c_d_id = $3 AND c_id = $4", - &[], + vec![Decimal(None, None), Smallint, Tinyint, Integer], + vec![], ), stmt( "INSERT INTO history(h_c_d_id, h_c_w_id, h_c_id, h_d_id, h_w_id, h_date, h_amount, h_data) VALUES($1, $2, $3, $4, $5, $6, $7, $8)", - &[], + vec![Tinyint, Smallint, Integer, Tinyint, Smallint, TimeStamp(Some(6), false), Decimal(None, None), Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![], ), ], vec![ // "SELECT count(c_id) FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_last = $3" stmt( "SELECT count(c_id) FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_last = $3", - &[ColumnType::Int32], + vec![Smallint, Tinyint, Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![Integer], ), // "SELECT c_balance, c_first, c_middle, c_last FROM customer WHERE ... ORDER BY c_first" stmt( "SELECT c_balance, c_first, c_middle, c_last FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_last = $3 ORDER BY c_first", - &[ - ColumnType::Decimal, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, + vec![Smallint, Tinyint, Varchar(None, kite_sql::types::CharLengthUnits::Characters)], + vec![ + Decimal(None, None), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), ], ), // "SELECT c_balance, c_first, c_middle, c_last FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_id = $3" stmt( "SELECT c_balance, c_first, c_middle, c_last FROM customer WHERE c_w_id = $1 AND c_d_id = $2 AND c_id = $3", - &[ - ColumnType::Decimal, - ColumnType::Utf8, - ColumnType::Utf8, - ColumnType::Utf8, + vec![Smallint, Tinyint, Integer], + vec![ + Decimal(None, None), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), + Varchar(None, kite_sql::types::CharLengthUnits::Characters), ], ), // "SELECT o_id, o_entry_d, COALESCE(o_carrier_id,0) FROM orders ..." stmt( "SELECT o_id, o_entry_d, COALESCE(o_carrier_id,0) FROM orders WHERE o_w_id = $1 AND o_d_id = $2 AND o_c_id = $3 AND o_id = (SELECT MAX(o_id) FROM orders WHERE o_w_id = $4 AND o_d_id = $5 AND o_c_id = $6)", - &[ColumnType::Int32, ColumnType::DateTime, ColumnType::Int32], + vec![Smallint, Tinyint, Integer, Smallint, Tinyint, Integer], + vec![Integer, DateTime, Integer], ), // "SELECT ol_i_id, ol_supply_w_id, ol_quantity, ol_amount, ol_delivery_d FROM order_line ..." stmt( "SELECT ol_i_id, ol_supply_w_id, ol_quantity, ol_amount, ol_delivery_d FROM order_line WHERE ol_w_id = $1 AND ol_d_id = $2 AND ol_o_id = $3", - &[ - ColumnType::Int32, - ColumnType::Int16, - ColumnType::Int8, - ColumnType::Decimal, - ColumnType::NullableDateTime, + vec![Smallint, Tinyint, Integer], + vec![ + Integer, + Smallint, + Tinyint, + Decimal(None, None), + DateTime, ], ), ], @@ -546,59 +575,81 @@ fn statement_specs() -> Vec> { // "SELECT COALESCE(MIN(no_o_id),0) FROM new_orders WHERE no_d_id = $1 AND no_w_id = $2" stmt( "SELECT COALESCE(MIN(no_o_id),0) FROM new_orders WHERE no_d_id = $1 AND no_w_id = $2", - &[ColumnType::Int32], + vec![Tinyint, Smallint], + vec![Integer], ), // "DELETE FROM new_orders WHERE no_o_id = $1 AND no_d_id = $2 AND no_w_id = $3" stmt( "DELETE FROM new_orders WHERE no_o_id = $1 AND no_d_id = $2 AND no_w_id = $3", - &[], + vec![Integer, Tinyint, Smallint], + vec![], ), // "SELECT o_c_id FROM orders WHERE o_id = $1 AND o_d_id = $2 AND o_w_id = $3" stmt( "SELECT o_c_id FROM orders WHERE o_id = $1 AND o_d_id = $2 AND o_w_id = $3", - &[ColumnType::Int32], + vec![Integer, Tinyint, Smallint], + vec![Integer], ), // "UPDATE orders SET o_carrier_id = $1 WHERE o_id = $2 AND o_d_id = $3 AND o_w_id = $4" stmt( "UPDATE orders SET o_carrier_id = $1 WHERE o_id = $2 AND o_d_id = $3 AND o_w_id = $4", - &[], + vec![Tinyint, Integer, Tinyint, Smallint], + vec![], ), // "UPDATE order_line SET ol_delivery_d = $1 WHERE ol_o_id = $2 AND ol_d_id = $3 AND ol_w_id = $4" stmt( "UPDATE order_line SET ol_delivery_d = $1 WHERE ol_o_id = $2 AND ol_d_id = $3 AND ol_w_id = $4", - &[], + vec![DateTime, Integer, Tinyint, Smallint], + vec![], ), // "SELECT SUM(ol_amount) FROM order_line WHERE ol_o_id = $1 AND ol_d_id = $2 AND ol_w_id = $3" stmt( "SELECT SUM(ol_amount) FROM order_line WHERE ol_o_id = $1 AND ol_d_id = $2 AND ol_w_id = $3", - &[ColumnType::Decimal], + vec![Integer, Tinyint, Smallint], + vec![Decimal(None, None)], ), // "UPDATE customer SET c_balance = c_balance + $1 , c_delivery_cnt = c_delivery_cnt + 1 WHERE c_id = $2 ..." stmt( "UPDATE customer SET c_balance = c_balance + $1 , c_delivery_cnt = c_delivery_cnt + 1 WHERE c_id = $2 AND c_d_id = $3 AND c_w_id = $4", - &[], + vec![Decimal(None, None), Integer, Tinyint, Smallint], + vec![], ), ], vec![ // "SELECT d_next_o_id FROM district WHERE d_id = $1 AND d_w_id = $2" stmt( "SELECT d_next_o_id FROM district WHERE d_id = $1 AND d_w_id = $2", - &[ColumnType::Int32], + vec![Tinyint, Smallint], + vec![Integer], ), stmt( "SELECT DISTINCT ol_i_id FROM order_line WHERE ol_w_id = $1 AND ol_d_id = $2 AND ol_o_id < $3 AND ol_o_id >= ($4 - 20)", - &[ColumnType::Int32], + vec![Smallint, Tinyint, Integer, Integer], + vec![Integer], ), stmt( "SELECT count(*) FROM stock WHERE s_w_id = $1 AND s_i_id = $2 AND s_quantity < $3", - &[ColumnType::Int32], + vec![Smallint, Integer, Smallint], + vec![Integer], ), ], ] } -fn stmt(sql: &'static str, result_types: &'static [ColumnType]) -> StatementSpec { - StatementSpec { sql, result_types } +fn stmt( + sql: &'static str, + parameter_types: Vec, + result_types: Vec, +) -> StatementSpec { + StatementSpec { + sql, + parameters: parameter_types + .into_iter() + .enumerate() + .map(|(i, ty)| (i + 1, ty)) + .collect(), + result_types, + } } fn print_summary_table(success: &[usize], late: &[usize], failure: &[usize], elapsed: Duration) { diff --git a/tpcc/src/new_ord.rs b/tpcc/src/new_ord.rs index 5c4376f3..ce2d44fd 100644 --- a/tpcc/src/new_ord.rs +++ b/tpcc/src/new_ord.rs @@ -88,10 +88,10 @@ impl TpccTransaction for NewOrd { tx.with_query_one( &mut statements[0], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int16(args.w_id as i16)), - ("$3", DataValue::Int8(args.d_id as i8)), - ("$4", DataValue::Int64(args.c_id as i64)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int16(args.w_id as i16)), + (3, DataValue::Int8(args.d_id as i8)), + (4, DataValue::Int64(args.c_id as i64)), ], &mut |tuple| { c_discount = tuple.values[0].decimal().unwrap(); @@ -111,9 +111,9 @@ impl TpccTransaction for NewOrd { tx.with_query_one( &mut statements[1], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int32(args.c_id as i32)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int32(args.c_id as i32)), ], &mut |tuple| { c_discount = tuple.values[0].decimal().unwrap(); @@ -126,7 +126,7 @@ impl TpccTransaction for NewOrd { let mut w_tax = Decimal::default(); tx.with_query_one( &mut statements[2], - &[("$1", DataValue::Int16(args.w_id as i16))], + &[(1, DataValue::Int16(args.w_id as i16))], &mut |tuple| { w_tax = tuple.values[0].decimal().unwrap(); Ok(()) @@ -141,8 +141,8 @@ impl TpccTransaction for NewOrd { tx.with_query_one( &mut statements[3], &[ - ("$1", DataValue::Int8(args.d_id as i8)), - ("$2", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int8(args.d_id as i8)), + (2, DataValue::Int16(args.w_id as i16)), ], &mut |tuple| { d_next_o_id = tuple.values[0].i32().unwrap(); @@ -154,9 +154,9 @@ impl TpccTransaction for NewOrd { tx.execute_drain( &mut statements[4], &[ - ("$1", DataValue::Int32(d_next_o_id)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int32(d_next_o_id)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), ], )?; let o_id = d_next_o_id; @@ -164,22 +164,22 @@ impl TpccTransaction for NewOrd { tx.execute_drain( &mut statements[5], &[ - ("$1", DataValue::Int32(o_id)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), - ("$4", DataValue::Int32(args.c_id as i32)), - ("$5", DataValue::from(&now)), - ("$6", DataValue::Int8(args.o_ol_cnt as i8)), - ("$7", DataValue::Int8(args.o_all_local as i8)), + (1, DataValue::Int32(o_id)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), + (4, DataValue::Int32(args.c_id as i32)), + (5, DataValue::from(&now)), + (6, DataValue::Int8(args.o_ol_cnt as i8)), + (7, DataValue::Int8(args.o_all_local as i8)), ], )?; // "INSERT INTO new_orders (no_o_id, no_d_id, no_w_id) VALUES (?,?,?)" tx.execute_drain( &mut statements[6], &[ - ("$1", DataValue::Int32(o_id)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int32(o_id)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), ], )?; let mut ol_num_seq = vec![0; MAX_NUM_ITEMS]; @@ -210,7 +210,7 @@ impl TpccTransaction for NewOrd { let ol_i_id = args.item_id[ol_num_seq[ol_number - 1]]; let ol_quantity = args.qty[ol_num_seq[ol_number - 1]]; // "SELECT i_price, i_name, i_data FROM item WHERE i_id = ?" - let params = [("$1", DataValue::Int32(ol_i_id as i32))]; + let params = [(1, DataValue::Int32(ol_i_id as i32))]; let mut i_price = Decimal::default(); let mut i_name = String::new(); let mut i_data = String::new(); @@ -226,8 +226,8 @@ impl TpccTransaction for NewOrd { // "SELECT s_quantity, s_data, s_dist_01, s_dist_02, s_dist_03, s_dist_04, s_dist_05, s_dist_06, s_dist_07, s_dist_08, s_dist_09, s_dist_10 FROM stock WHERE s_i_id = ? AND s_w_id = ? FOR UPDATE" let params = [ - ("$1", DataValue::Int32(ol_i_id as i32)), - ("$2", DataValue::Int16(ol_supply_w_id as i16)), + (1, DataValue::Int32(ol_i_id as i32)), + (2, DataValue::Int16(ol_supply_w_id as i16)), ]; let mut s_quantity = 0; let mut s_data = String::new(); @@ -276,9 +276,9 @@ impl TpccTransaction for NewOrd { }; // "UPDATE stock SET s_quantity = ? WHERE s_i_id = ? AND s_w_id = ?" let params = [ - ("$1", DataValue::Int16(s_quantity)), - ("$2", DataValue::Int32(ol_i_id as i32)), - ("$3", DataValue::Int16(ol_supply_w_id as i16)), + (1, DataValue::Int16(s_quantity)), + (2, DataValue::Int32(ol_i_id as i32)), + (3, DataValue::Int16(ol_supply_w_id as i16)), ]; tx.execute_drain(&mut statements[9], ¶ms)?; @@ -294,15 +294,15 @@ impl TpccTransaction for NewOrd { amt[ol_num_seq[ol_number - 1]] = ol_amount; // "INSERT INTO order_line (ol_o_id, ol_d_id, ol_w_id, ol_number, ol_i_id, ol_supply_w_id, ol_quantity, ol_amount, ol_dist_info) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)" let params = [ - ("$1", DataValue::Int32(o_id)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int16(args.w_id as i16)), - ("$4", DataValue::Int8(ol_number as i8)), - ("$5", DataValue::Int32(ol_i_id as i32)), - ("$6", DataValue::Int16(ol_supply_w_id as i16)), - ("$7", DataValue::Int8(ol_quantity as i8)), - ("$8", DataValue::Decimal(ol_amount.round_dp(2))), - ("$9", DataValue::from(ol_dist_info)), + (1, DataValue::Int32(o_id)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int16(args.w_id as i16)), + (4, DataValue::Int8(ol_number as i8)), + (5, DataValue::Int32(ol_i_id as i32)), + (6, DataValue::Int16(ol_supply_w_id as i16)), + (7, DataValue::Int8(ol_quantity as i8)), + (8, DataValue::Decimal(ol_amount.round_dp(2))), + (9, DataValue::from(ol_dist_info)), ]; tx.execute_drain(&mut statements[10], ¶ms)?; } diff --git a/tpcc/src/order_stat.rs b/tpcc/src/order_stat.rs index fa76ba1d..5243c599 100644 --- a/tpcc/src/order_stat.rs +++ b/tpcc/src/order_stat.rs @@ -63,9 +63,9 @@ impl TpccTransaction for OrderStat { tx.with_query_one( &mut statements[0], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::from(args.c_last.clone())), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::from(args.c_last.clone())), ], &mut |tuple| { name_cnt = tuple.values[0].i32().unwrap() as usize; @@ -74,9 +74,9 @@ impl TpccTransaction for OrderStat { )?; // "SELECT c_balance, c_first, c_middle, c_last FROM customer WHERE c_w_id = ? AND c_d_id = ? AND c_last = ? ORDER BY c_first" let params = [ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::from(args.c_last.clone())), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::from(args.c_last.clone())), ]; if name_cnt % 2 == 1 { name_cnt += 1; @@ -103,9 +103,9 @@ impl TpccTransaction for OrderStat { tx.with_query_one( &mut statements[2], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int32(args.c_id as i32)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int32(args.c_id as i32)), ], &mut |tuple| { c_balance = tuple.values[0].decimal().unwrap(); @@ -119,12 +119,12 @@ impl TpccTransaction for OrderStat { }; // "SELECT o_id, o_entry_d, COALESCE(o_carrier_id,0) FROM orders WHERE o_w_id = ? AND o_d_id = ? AND o_c_id = ? AND o_id = (SELECT MAX(o_id) FROM orders WHERE o_w_id = ? AND o_d_id = ? AND o_c_id = ?)" let params = [ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int32(args.c_id as i32)), - ("$4", DataValue::Int16(args.w_id as i16)), - ("$5", DataValue::Int8(args.d_id as i8)), - ("$6", DataValue::Int32(args.c_id as i32)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int32(args.c_id as i32)), + (4, DataValue::Int16(args.w_id as i16)), + (5, DataValue::Int8(args.d_id as i8)), + (6, DataValue::Int32(args.c_id as i32)), ]; let mut o_id = 0; tx.with_query_one(&mut statements[3], ¶ms, &mut |tuple| { @@ -133,9 +133,9 @@ impl TpccTransaction for OrderStat { })?; // "SELECT ol_i_id, ol_supply_w_id, ol_quantity, ol_amount, ol_delivery_d FROM order_line WHERE ol_w_id = ? AND ol_d_id = ? AND ol_o_id = ?" let params = [ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int32(o_id)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int32(o_id)), ]; tx.with_query_one(&mut statements[4], ¶ms, &mut |_| Ok(()))?; // let ol_i_id = tuple.values[0].i32(); diff --git a/tpcc/src/payment.rs b/tpcc/src/payment.rs index 5b4f846a..80ebe1fd 100644 --- a/tpcc/src/payment.rs +++ b/tpcc/src/payment.rs @@ -75,8 +75,8 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[0], &[ - ("$1", DataValue::Decimal(args.h_amount)), - ("$2", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Decimal(args.h_amount)), + (2, DataValue::Int16(args.w_id as i16)), ], )?; // "SELECT w_street_1, w_street_2, w_city, w_state, w_zip, w_name FROM warehouse WHERE w_id = ?" @@ -88,7 +88,7 @@ impl TpccTransaction for Payment { let mut w_name = String::new(); tx.with_query_one( &mut statements[1], - &[("$1", DataValue::Int16(args.w_id as i16))], + &[(1, DataValue::Int16(args.w_id as i16))], &mut |tuple| { w_street_1 = tuple.values[0].utf8().unwrap().to_string(); w_street_2 = tuple.values[1].utf8().unwrap().to_string(); @@ -104,9 +104,9 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[2], &[ - ("$1", DataValue::Decimal(args.h_amount)), - ("$2", DataValue::Int16(args.w_id as i16)), - ("$3", DataValue::Int8(args.d_id as i8)), + (1, DataValue::Decimal(args.h_amount)), + (2, DataValue::Int16(args.w_id as i16)), + (3, DataValue::Int8(args.d_id as i8)), ], )?; @@ -120,8 +120,8 @@ impl TpccTransaction for Payment { tx.with_query_one( &mut statements[3], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), ], &mut |tuple| { d_street_1 = tuple.values[0].utf8().unwrap().to_string(); @@ -141,9 +141,9 @@ impl TpccTransaction for Payment { tx.with_query_one( &mut statements[4], &[ - ("$1", DataValue::Int16(args.c_w_id as i16)), - ("$2", DataValue::Int8(args.c_d_id as i8)), - ("$3", DataValue::from(args.c_last.clone())), + (1, DataValue::Int16(args.c_w_id as i16)), + (2, DataValue::Int8(args.c_d_id as i8)), + (3, DataValue::from(args.c_last.clone())), ], &mut |tuple| { name_cnt = tuple.values[0].i32().unwrap(); @@ -152,9 +152,9 @@ impl TpccTransaction for Payment { )?; // "SELECT c_id FROM customer WHERE c_w_id = ? AND c_d_id = ? AND c_last = ? ORDER BY c_first" let params = [ - ("$1", DataValue::Int16(args.c_w_id as i16)), - ("$2", DataValue::Int8(args.c_d_id as i8)), - ("$3", DataValue::from(args.c_last.clone())), + (1, DataValue::Int16(args.c_w_id as i16)), + (2, DataValue::Int8(args.c_d_id as i8)), + (3, DataValue::from(args.c_last.clone())), ]; if name_cnt % 2 == 1 { name_cnt += 1; @@ -183,9 +183,9 @@ impl TpccTransaction for Payment { tx.with_query_one( &mut statements[6], &[ - ("$1", DataValue::Int16(args.c_w_id as i16)), - ("$2", DataValue::Int8(args.c_d_id as i8)), - ("$3", DataValue::Int32(c_id)), + (1, DataValue::Int16(args.c_w_id as i16)), + (2, DataValue::Int8(args.c_d_id as i8)), + (3, DataValue::Int32(c_id)), ], &mut |tuple| { c_first = tuple.values[0].utf8().unwrap().to_string(); @@ -214,9 +214,9 @@ impl TpccTransaction for Payment { tx.with_query_one( &mut statements[7], &[ - ("$1", DataValue::Int16(args.c_w_id as i16)), - ("$2", DataValue::Int8(args.c_d_id as i8)), - ("$3", DataValue::Int32(c_id)), + (1, DataValue::Int16(args.c_w_id as i16)), + (2, DataValue::Int8(args.c_d_id as i8)), + (3, DataValue::Int32(c_id)), ], &mut |tuple| { c_data = tuple.values[0].utf8().unwrap().to_string(); @@ -231,11 +231,11 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[8], &[ - ("$1", DataValue::Decimal(c_balance)), - ("$2", DataValue::from(c_data)), - ("$3", DataValue::Int16(args.c_w_id as i16)), - ("$4", DataValue::Int8(args.c_d_id as i8)), - ("$5", DataValue::Int32(c_id)), + (1, DataValue::Decimal(c_balance)), + (2, DataValue::from(c_data)), + (3, DataValue::Int16(args.c_w_id as i16)), + (4, DataValue::Int8(args.c_d_id as i8)), + (5, DataValue::Int32(c_id)), ], )?; } else { @@ -243,10 +243,10 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[9], &[ - ("$1", DataValue::Decimal(c_balance)), - ("$2", DataValue::Int16(args.c_w_id as i16)), - ("$3", DataValue::Int8(args.c_d_id as i8)), - ("$4", DataValue::Int32(c_id)), + (1, DataValue::Decimal(c_balance)), + (2, DataValue::Int16(args.c_w_id as i16)), + (3, DataValue::Int8(args.c_d_id as i8)), + (4, DataValue::Int32(c_id)), ], )?; } @@ -255,10 +255,10 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[9], &[ - ("$1", DataValue::Decimal(c_balance)), - ("$2", DataValue::Int16(args.c_w_id as i16)), - ("$3", DataValue::Int8(args.c_d_id as i8)), - ("$4", DataValue::Int32(c_id)), + (1, DataValue::Decimal(c_balance)), + (2, DataValue::Int16(args.c_w_id as i16)), + (3, DataValue::Int8(args.c_d_id as i8)), + (4, DataValue::Int32(c_id)), ], )?; } @@ -267,14 +267,14 @@ impl TpccTransaction for Payment { tx.execute_drain( &mut statements[10], &[ - ("$1", DataValue::Int8(args.c_d_id as i8)), - ("$2", DataValue::Int16(args.c_w_id as i16)), - ("$3", DataValue::Int32(c_id)), - ("$4", DataValue::Int8(args.d_id as i8)), - ("$5", DataValue::Int16(args.w_id as i16)), - ("$6", DataValue::Time64(now.timestamp_micros(), 6, false)), - ("$7", DataValue::Decimal(args.h_amount)), - ("$8", DataValue::from(h_data)), + (1, DataValue::Int8(args.c_d_id as i8)), + (2, DataValue::Int16(args.c_w_id as i16)), + (3, DataValue::Int32(c_id)), + (4, DataValue::Int8(args.d_id as i8)), + (5, DataValue::Int16(args.w_id as i16)), + (6, DataValue::Time64(now.timestamp_micros(), 6, false)), + (7, DataValue::Decimal(args.h_amount)), + (8, DataValue::from(h_data)), ], )?; diff --git a/tpcc/src/slev.rs b/tpcc/src/slev.rs index b3942b17..5fc5a773 100644 --- a/tpcc/src/slev.rs +++ b/tpcc/src/slev.rs @@ -48,8 +48,8 @@ impl TpccTransaction for Slev { tx.with_query_one( &mut statements[0], &[ - ("$1", DataValue::Int8(args.d_id as i8)), - ("$2", DataValue::Int16(args.w_id as i16)), + (1, DataValue::Int8(args.d_id as i8)), + (2, DataValue::Int16(args.w_id as i16)), ], &mut |tuple| { d_next_o_id = tuple.values[0].i32().unwrap(); @@ -61,10 +61,10 @@ impl TpccTransaction for Slev { tx.with_query_all( &mut statements[1], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int8(args.d_id as i8)), - ("$3", DataValue::Int32(d_next_o_id)), - ("$4", DataValue::Int32(d_next_o_id)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int8(args.d_id as i8)), + (3, DataValue::Int32(d_next_o_id)), + (4, DataValue::Int32(d_next_o_id)), ], &mut |tuple| { item_ids.push(tuple.values[0].i32().unwrap()); @@ -77,9 +77,9 @@ impl TpccTransaction for Slev { tx.with_query_one( &mut statements[2], &[ - ("$1", DataValue::Int16(args.w_id as i16)), - ("$2", DataValue::Int32(item_id)), - ("$3", DataValue::Int16(args.level as i16)), + (1, DataValue::Int16(args.w_id as i16)), + (2, DataValue::Int32(item_id)), + (3, DataValue::Int16(args.level as i16)), ], &mut |tuple| { _low_stock += tuple.values[0].i32().unwrap();