diff --git a/dozer-sql/expression/src/builder.rs b/dozer-sql/expression/src/builder.rs index cefcda0039..06b6178e44 100644 --- a/dozer-sql/expression/src/builder.rs +++ b/dozer-sql/expression/src/builder.rs @@ -8,7 +8,7 @@ use dozer_types::models::udf_config::{UdfConfig, UdfType}; use dozer_types::types::FieldType; use dozer_types::{ ordered_float::OrderedFloat, - types::{Field, FieldDefinition, Schema, SourceDefinition}, + types::{Field, Schema, SourceDefinition}, }; use sqlparser::ast::{ BinaryOperator as SqlBinaryOperator, DataType, DateTimeField, Expr as SqlExpr, Expr, Function, @@ -211,63 +211,27 @@ impl ExpressionBuilder { } }; - let matching_by_field: Vec<(usize, &FieldDefinition)> = schema - .fields - .iter() - .enumerate() - .filter(|(_idx, f)| &f.name == src_field) - .collect(); - - match matching_by_field.len() { - 1 => Ok(Expression::Column { - index: matching_by_field[0].0, - }), - _ => match src_table_or_alias { - None => Err(Error::InvalidIdent(ident.to_vec())), - Some(src_table_or_alias) => { - let matching_by_table_or_alias: Vec<(usize, &FieldDefinition)> = - matching_by_field - .into_iter() - .filter(|(_idx, field)| match &field.source { - SourceDefinition::Alias { name } => name == src_table_or_alias, - SourceDefinition::Table { - name, - connection: _, - } => name == src_table_or_alias, - _ => false, - }) - .collect(); - - match matching_by_table_or_alias.len() { - 1 => Ok(Expression::Column { - index: matching_by_table_or_alias[0].0, - }), - _ => match src_connection { - None => Err(Error::InvalidIdent(ident.to_vec())), - Some(src_connection) => { - let matching_by_connection: Vec<(usize, &FieldDefinition)> = - matching_by_table_or_alias - .into_iter() - .filter(|(_idx, field)| match &field.source { - SourceDefinition::Table { - name: _, - connection, - } => connection == src_connection, - _ => false, - }) - .collect(); - - match matching_by_connection.len() { - 1 => Ok(Expression::Column { - index: matching_by_connection[0].0, - }), - _ => Err(Error::InvalidIdent(ident.to_vec())), - } - } - }, - } - } - }, + let mut matching = schema.fields.iter().enumerate().filter(|(_, field)| { + if &field.name != src_field { + return false; + } + match (src_table_or_alias, src_connection, &field.source) { + (None, None, _) => true, + (Some(table), None, SourceDefinition::Alias { name }) => name == table, + ( + Some(table), + connection, + SourceDefinition::Table { + name, + connection: source, + }, + ) => name == table && connection.map_or(true, |connection| source == connection), + _ => false, + } + }); + match (matching.next(), matching.next()) { + (Some((index, _)), None) => Ok(Expression::Column { index }), + _ => Err(Error::InvalidIdent(ident.to_vec())), } } @@ -1062,3 +1026,104 @@ pub fn extend_schema_source_def(schema: &Schema, name: &NameOrAlias) -> Schema { output_schema } + +#[cfg(test)] +mod tests { + use super::*; + use dozer_types::types::FieldDefinition; + + fn resolve(schema: &Schema, name: &str) -> Result { + let ident = name.split('.').map(Ident::new).collect::>(); + ExpressionBuilder::parse_sql_column(&ident, schema) + } + + fn schema_with_sources(sources: Vec) -> Schema { + Schema { + fields: sources + .into_iter() + .map(|source| FieldDefinition::new("id".into(), FieldType::Int, false, source)) + .collect(), + primary_index: vec![], + } + } + + #[test] + fn unique_column_still_requires_matching_qualifiers() { + let schema = schema_with_sources(vec![SourceDefinition::Table { + connection: "source".into(), + name: "inner_rows".into(), + }]); + for name in ["id", "inner_rows.id", "source.inner_rows.id"] { + assert_eq!( + resolve(&schema, name).unwrap(), + Expression::Column { index: 0 } + ); + } + for name in [ + "outer_rows.id", + "other.inner_rows.id", + "source.outer_rows.id", + "extra.source.inner_rows.id", + ] { + assert!( + matches!(resolve(&schema, name), Err(Error::InvalidIdent(_))), + "{name}" + ); + } + } + + #[test] + fn aliases_and_unqualified_sources_do_not_accept_outer_qualifiers() { + let alias = schema_with_sources(vec![SourceDefinition::Alias { + name: "inner_alias".into(), + }]); + for name in ["id", "inner_alias.id"] { + assert_eq!( + resolve(&alias, name).unwrap(), + Expression::Column { index: 0 } + ); + } + for name in ["outer_alias.id", "source.inner_alias.id"] { + assert!(matches!(resolve(&alias, name), Err(Error::InvalidIdent(_)))); + } + let dynamic = schema_with_sources(vec![SourceDefinition::Dynamic]); + assert_eq!( + resolve(&dynamic, "id").unwrap(), + Expression::Column { index: 0 } + ); + assert!(matches!( + resolve(&dynamic, "outer_rows.id"), + Err(Error::InvalidIdent(_)) + )); + } + + #[test] + fn all_qualifiers_are_applied_before_checking_ambiguity() { + let schema = schema_with_sources(vec![ + SourceDefinition::Table { + connection: "first".into(), + name: "rows".into(), + }, + SourceDefinition::Table { + connection: "second".into(), + name: "rows".into(), + }, + SourceDefinition::Table { + connection: "first".into(), + name: "other".into(), + }, + ]); + for name in ["id", "rows.id", "third.rows.id", "second.other.id"] { + assert!( + matches!(resolve(&schema, name), Err(Error::InvalidIdent(_))), + "{name}" + ); + } + for (name, index) in [("first.rows.id", 0), ("second.rows.id", 1), ("other.id", 2)] { + assert_eq!( + resolve(&schema, name).unwrap(), + Expression::Column { index } + ); + } + } +} diff --git a/dozer-sql/expression/src/execution.rs b/dozer-sql/expression/src/execution.rs index 1c2891b234..ee35beac2a 100644 --- a/dozer-sql/expression/src/execution.rs +++ b/dozer-sql/expression/src/execution.rs @@ -544,7 +544,7 @@ fn get_binary_operator_type( match (left_field_type.return_type, right_field_type.return_type) { (FieldType::Boolean, FieldType::Boolean) => Ok(ExpressionType::new( FieldType::Boolean, - false, + left_field_type.nullable || right_field_type.nullable, SourceDefinition::Dynamic, false, )), @@ -558,7 +558,7 @@ fn get_binary_operator_type( | FieldType::Text, ) => Ok(ExpressionType::new( FieldType::Boolean, - false, + left_field_type.nullable || right_field_type.nullable, SourceDefinition::Dynamic, false, )), @@ -572,7 +572,7 @@ fn get_binary_operator_type( FieldType::Boolean, ) => Ok(ExpressionType::new( FieldType::Boolean, - false, + left_field_type.nullable || right_field_type.nullable, SourceDefinition::Dynamic, false, )), diff --git a/dozer-sql/expression/src/logical.rs b/dozer-sql/expression/src/logical.rs index 63a9017840..6474558174 100644 --- a/dozer-sql/expression/src/logical.rs +++ b/dozer-sql/expression/src/logical.rs @@ -10,66 +10,13 @@ pub fn evaluate_and( right: &mut Expression, record: &Record, ) -> Result { - let l_field = left.evaluate(record, schema)?; - let r_field = right.evaluate(record, schema)?; - match l_field { - Field::Boolean(true) => match r_field { - Field::Boolean(true) => Ok(Field::Boolean(true)), - Field::Boolean(false) => Ok(Field::Boolean(false)), - Field::Null => Ok(Field::Boolean(false)), - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(r_field, "AND".to_string())), - }, - Field::Boolean(false) => match r_field { - Field::Boolean(true) => Ok(Field::Boolean(false)), - Field::Boolean(false) => Ok(Field::Boolean(false)), - Field::Null => Ok(Field::Boolean(false)), - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(r_field, "AND".to_string())), - }, - Field::Null => Ok(Field::Boolean(false)), - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(l_field, "AND".to_string())), - } + let left = boolean_operand(left.evaluate(record, schema)?, "AND")?; + let right = boolean_operand(right.evaluate(record, schema)?, "AND")?; + Ok(match (left, right) { + (Some(false), _) | (_, Some(false)) => Field::Boolean(false), + (Some(true), Some(true)) => Field::Boolean(true), + _ => Field::Null, + }) } pub fn evaluate_or( @@ -78,64 +25,20 @@ pub fn evaluate_or( right: &mut Expression, record: &Record, ) -> Result { - let l_field = left.evaluate(record, schema)?; - let r_field = right.evaluate(record, schema)?; - match l_field { - Field::Boolean(true) => match r_field { - Field::Boolean(false) => Ok(Field::Boolean(true)), - Field::Boolean(true) => Ok(Field::Boolean(true)), - Field::Null => Ok(Field::Boolean(true)), - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(r_field, "OR".to_string())), - }, - Field::Boolean(false) | Field::Null => match right.evaluate(record, schema)? { - Field::Boolean(false) => Ok(Field::Boolean(false)), - Field::Boolean(true) => Ok(Field::Boolean(true)), - Field::Null => Ok(Field::Boolean(false)), - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(r_field, "OR".to_string())), - }, - Field::UInt(_) - | Field::U128(_) - | Field::Int(_) - | Field::Int8(_) - | Field::I128(_) - | Field::Float(_) - | Field::String(_) - | Field::Text(_) - | Field::Binary(_) - | Field::Decimal(_) - | Field::Timestamp(_) - | Field::Date(_) - | Field::Json(_) - | Field::Point(_) - | Field::Duration(_) => Err(Error::InvalidType(l_field, "OR".to_string())), + let left = boolean_operand(left.evaluate(record, schema)?, "OR")?; + let right = boolean_operand(right.evaluate(record, schema)?, "OR")?; + Ok(match (left, right) { + (Some(true), _) | (_, Some(true)) => Field::Boolean(true), + (Some(false), Some(false)) => Field::Boolean(false), + _ => Field::Null, + }) +} + +fn boolean_operand(value: Field, operator: &str) -> Result, Error> { + match value { + Field::Boolean(value) => Ok(Some(value)), + Field::Null => Ok(None), + value => Err(Error::InvalidType(value, operator.to_owned())), } } @@ -171,12 +74,129 @@ pub fn evaluate_not( mod tests { use super::*; + use crate::operator::{BinaryOperatorType, UnaryOperatorType}; use dozer_types::types::Record; - use dozer_types::types::{Field, Schema}; + use dozer_types::types::{Field, FieldDefinition, FieldType, Schema, SourceDefinition}; use dozer_types::{ordered_float::OrderedFloat, rust_decimal::Decimal}; use proptest::prelude::*; use Expression::Literal; + #[test] + fn logical_truth_tables_preserve_unknown() { + use Field::{Boolean, Null}; + let cases = [ + (Boolean(true), Boolean(true), Boolean(true), Boolean(true)), + (Boolean(true), Boolean(false), Boolean(false), Boolean(true)), + (Boolean(true), Null, Null, Boolean(true)), + (Boolean(false), Boolean(true), Boolean(false), Boolean(true)), + ( + Boolean(false), + Boolean(false), + Boolean(false), + Boolean(false), + ), + (Boolean(false), Null, Boolean(false), Null), + (Null, Boolean(true), Null, Boolean(true)), + (Null, Boolean(false), Boolean(false), Null), + (Null, Null, Null, Null), + ]; + let schema = Schema::default(); + let record = Record::new(vec![]); + for (left, right, expected_and, expected_or) in cases { + let mut left = Literal(left); + let mut right = Literal(right); + assert_eq!( + evaluate_and(&schema, &mut left, &mut right, &record).unwrap(), + expected_and + ); + assert_eq!( + evaluate_or(&schema, &mut left, &mut right, &record).unwrap(), + expected_or + ); + } + } + + #[test] + fn negating_a_compound_unknown_does_not_select_it() { + let schema = Schema { + fields: vec![FieldDefinition::new( + "membership".into(), + FieldType::Boolean, + true, + SourceDefinition::Dynamic, + )], + primary_index: vec![], + }; + let record = Record::new(vec![Field::Null]); + for (operator, other) in [ + (BinaryOperatorType::And, true), + (BinaryOperatorType::Or, false), + ] { + let mut expression = Expression::UnaryOperator { + operator: UnaryOperatorType::Not, + arg: Box::new(Expression::BinaryOperator { + left: Box::new(Expression::Column { index: 0 }), + operator, + right: Box::new(Literal(Field::Boolean(other))), + }), + }; + assert_eq!(expression.evaluate(&record, &schema).unwrap(), Field::Null); + assert!(expression.get_type(&schema).unwrap().nullable); + } + } + + #[test] + fn logical_result_nullability_includes_both_operands() { + for (left_nullable, right_nullable) in + [(false, false), (false, true), (true, false), (true, true)] + { + let schema = Schema { + fields: vec![ + FieldDefinition::new( + "left".into(), + FieldType::Boolean, + left_nullable, + SourceDefinition::Dynamic, + ), + FieldDefinition::new( + "right".into(), + FieldType::Boolean, + right_nullable, + SourceDefinition::Dynamic, + ), + ], + primary_index: vec![], + }; + for operator in [BinaryOperatorType::And, BinaryOperatorType::Or] { + let expression = Expression::BinaryOperator { + left: Box::new(Expression::Column { index: 0 }), + operator, + right: Box::new(Expression::Column { index: 1 }), + }; + let result_type = expression.get_type(&schema).unwrap(); + assert_eq!(result_type.return_type, FieldType::Boolean); + assert_eq!(result_type.nullable, left_nullable || right_nullable); + } + } + } + + #[test] + fn unknown_does_not_hide_an_invalid_operand_type() { + let record = Record::new(vec![]); + for (left, right) in [(Field::Null, Field::Int(1)), (Field::Int(1), Field::Null)] { + let mut left = Literal(left); + let mut right = Literal(right); + assert!(matches!( + evaluate_and(&Schema::default(), &mut left, &mut right, &record), + Err(Error::InvalidType(Field::Int(1), _)) + )); + assert!(matches!( + evaluate_or(&Schema::default(), &mut left, &mut right, &record), + Err(Error::InvalidType(Field::Int(1), _)) + )); + } + } + #[test] fn test_logical() { proptest!( @@ -226,65 +246,72 @@ mod tests { let row = Record::new(vec![]); let mut l = Box::new(Literal(Field::Boolean(bool1))); let mut r = Box::new(Literal(Field::Boolean(bool2))); - assert!(matches!( - evaluate_and(&Schema::default(), &mut l, &mut r, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(_ans) - )); + assert_eq!( + evaluate_and(&Schema::default(), &mut l, &mut r, &row).unwrap(), + Field::Boolean(bool1 && bool2) + ); } fn _test_bool_null_and(f1: Field, f2: Field) { let row = Record::new(vec![]); + let expected = if f1 == Field::Boolean(false) || f2 == Field::Boolean(false) { + Field::Boolean(false) + } else { + Field::Null + }; let mut l = Box::new(Literal(f1)); let mut r = Box::new(Literal(f2)); - assert!(matches!( - evaluate_and(&Schema::default(), &mut l, &mut r, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(false) - )); + assert_eq!( + evaluate_and(&Schema::default(), &mut l, &mut r, &row).unwrap(), + expected + ); } fn _test_bool_bool_or(bool1: bool, bool2: bool) { let row = Record::new(vec![]); let mut l = Box::new(Literal(Field::Boolean(bool1))); let mut r = Box::new(Literal(Field::Boolean(bool2))); - assert!(matches!( - evaluate_or(&Schema::default(), &mut l, &mut r, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(_ans) - )); + assert_eq!( + evaluate_or(&Schema::default(), &mut l, &mut r, &row).unwrap(), + Field::Boolean(bool1 || bool2) + ); } fn _test_bool_null_or(_bool: bool) { let row = Record::new(vec![]); let mut l = Box::new(Literal(Field::Boolean(_bool))); let mut r = Box::new(Literal(Field::Null)); - assert!(matches!( - evaluate_or(&Schema::default(), &mut l, &mut r, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(_bool) - )); + assert_eq!( + evaluate_or(&Schema::default(), &mut l, &mut r, &row).unwrap(), + if _bool { + Field::Boolean(true) + } else { + Field::Null + } + ); } fn _test_null_bool_or(_bool: bool) { let row = Record::new(vec![]); let mut l = Box::new(Literal(Field::Null)); let mut r = Box::new(Literal(Field::Boolean(_bool))); - assert!(matches!( - evaluate_or(&Schema::default(), &mut l, &mut r, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(_bool) - )); + assert_eq!( + evaluate_or(&Schema::default(), &mut l, &mut r, &row).unwrap(), + if _bool { + Field::Boolean(true) + } else { + Field::Null + } + ); } fn _test_bool_not(bool: bool) { let row = Record::new(vec![]); let mut v = Box::new(Literal(Field::Boolean(bool))); - assert!(matches!( - evaluate_not(&Schema::default(), &mut v, &row) - .unwrap_or_else(|e| panic!("{}", e.to_string())), - Field::Boolean(_ans) - )); + assert_eq!( + evaluate_not(&Schema::default(), &mut v, &row).unwrap(), + Field::Boolean(!bool) + ); } fn _test_bool_non_bool_and(f1: Field, f2: Field) { diff --git a/dozer-sql/src/aggregation/factory.rs b/dozer-sql/src/aggregation/factory.rs index c91d333b0b..0acc3f3367 100644 --- a/dozer-sql/src/aggregation/factory.rs +++ b/dozer-sql/src/aggregation/factory.rs @@ -1,11 +1,13 @@ use crate::planner::projection::CommonPlanner; use crate::projection::processor::ProjectionProcessor; +use crate::selection::replay::ReplayRequirement; use crate::{aggregation::processor::AggregationProcessor, errors::PipelineError}; use dozer_core::event::EventHub; use dozer_core::{ node::{PortHandle, Processor, ProcessorFactory}, DEFAULT_PORT_HANDLE, }; +use dozer_sql_expression::execution::Expression; use dozer_sql_expression::sqlparser::ast::{Expr, SelectItem}; use dozer_types::errors::internal::BoxedError; use dozer_types::models::udf_config::UdfConfig; @@ -25,6 +27,7 @@ pub struct AggregationProcessorFactory { enable_probabilistic_optimizations: bool, udfs: Vec, runtime: Arc, + replay: ReplayRequirement, /// Type name can only be determined after schema propagation. type_name: Mutex>, @@ -48,10 +51,15 @@ impl AggregationProcessorFactory { enable_probabilistic_optimizations, udfs, runtime, + replay: ReplayRequirement::default(), type_name: Mutex::new(None), } } + pub(crate) fn replay_requirement(&self) -> ReplayRequirement { + self.replay.clone() + } + async fn get_planner(&self, input_schema: Schema) -> Result { let mut projection_planner = CommonPlanner::new(input_schema, self.udfs.as_slice(), self.runtime.clone()); @@ -62,6 +70,33 @@ impl AggregationProcessorFactory { self.having.clone(), ) .await?; + if self.replay.is_required() + && !projection_planner.aggregation_output.is_empty() + && projection_planner.groupby.is_empty() + { + // The existing aggregate processor emits no row for empty input; + // global COUNT/SUM cannot supply SQL's empty-set row to membership. + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support aggregates without GROUP BY".into(), + )); + } + for expression in projection_planner + .projection_output + .iter() + .chain(projection_planner.groupby.iter()) + .chain(projection_planner.having.iter()) + { + self.replay.validate(expression)?; + } + for aggregate in &projection_planner.aggregation_output { + if let Expression::AggregateFunction { args, .. } = aggregate { + for argument in args { + self.replay.validate(argument)?; + } + } else { + self.replay.validate(aggregate)?; + } + } Ok(projection_planner) } } diff --git a/dozer-sql/src/builder/from.rs b/dozer-sql/src/builder/from.rs index 6fa98bea03..c2720b8380 100644 --- a/dozer-sql/src/builder/from.rs +++ b/dozer-sql/src/builder/from.rs @@ -69,8 +69,17 @@ fn insert_table_processor_to_pipeline( pipeline_idx: usize, query_context: &mut QueryContext, ) -> Result { - let relation_name_or_alias = get_name_or_alias(&relation)?; + let mut relation_name_or_alias = get_name_or_alias(&relation)?; let product_input_name = get_from_source(relation, pipeline, query_context, pipeline_idx)?.0; + if relation_name_or_alias.1.is_none() + && query_context + .pipeline_map + .contains_key(&(pipeline_idx, product_input_name.clone())) + { + // A CTE exposes its own relation name, not the source table names which + // happened to produce its columns. + relation_name_or_alias.1 = Some(relation_name_or_alias.0.clone()); + } let processor_name = format!( "from:{}--{}", diff --git a/dozer-sql/src/builder/join.rs b/dozer-sql/src/builder/join.rs index 653010d927..c33b01c606 100644 --- a/dozer-sql/src/builder/join.rs +++ b/dozer-sql/src/builder/join.rs @@ -36,10 +36,16 @@ pub fn insert_join_to_pipeline( let mut left_name_or_alias = Some(get_name_or_alias(&left_table)?); let mut left_join_source = insert_join_source_to_pipeline(left_table, pipeline, pipeline_idx, query_context)?; + bind_relation_name( + &mut left_name_or_alias, + &left_join_source, + query_context, + pipeline_idx, + ); for join in from.joins { let right_table = join.relation; - let right_name_or_alias = Some(get_name_or_alias(&right_table)?); + let mut right_name_or_alias = Some(get_name_or_alias(&right_table)?); let right_join_source = insert_join_source_to_pipeline( right_table.clone(), pipeline, @@ -47,6 +53,13 @@ pub fn insert_join_to_pipeline( query_context, )?; + bind_relation_name( + &mut right_name_or_alias, + &right_join_source, + query_context, + pipeline_idx, + ); + let join_processor_name = format!("join_{}", query_context.get_next_processor_id()); if !query_context .processors_list @@ -103,6 +116,23 @@ pub fn insert_join_to_pipeline( } } +fn bind_relation_name( + name: &mut Option, + source: &JoinSource, + context: &QueryContext, + pipeline_idx: usize, +) { + if let (Some(name), JoinSource::Table(source)) = (name, source) { + if name.1.is_none() + && context + .pipeline_map + .contains_key(&(pipeline_idx, source.clone())) + { + name.1 = Some(name.0.clone()); + } + } +} + // TODO: refactor this fn insert_join_source_to_pipeline( source: TableFactor, diff --git a/dozer-sql/src/builder/membership.rs b/dozer-sql/src/builder/membership.rs new file mode 100644 index 0000000000..5ca22b72c7 --- /dev/null +++ b/dozer-sql/src/builder/membership.rs @@ -0,0 +1,371 @@ +use dozer_core::{app::AppPipeline, DEFAULT_PORT_HANDLE}; +use dozer_sql_expression::{ + builder::NameOrAlias, + sqlparser::ast::{ + Expr, FunctionArg, FunctionArgExpr, Ident, Query, SelectItem, SetExpr, SetOperator, + TableFactor, + }, +}; + +use crate::{ + errors::PipelineError, + selection::membership::{MembershipProcessorFactory, OUTER_PORT, VALUES_PORT}, +}; + +use super::{query_to_pipeline, OutputNodeInfo, QueryContext, TableInfo}; + +pub(super) fn plan( + predicate: &mut Expr, + input: OutputNodeInfo, + pipeline: &mut AppPipeline, + context: &mut QueryContext, + pipeline_idx: usize, +) -> Result<(OutputNodeInfo, Vec), PipelineError> { + let mut planner = MembershipPlanner { + pipeline, + context, + pipeline_idx, + input, + columns: Vec::new(), + }; + planner.rewrite(predicate)?; + Ok((planner.input, planner.columns)) +} + +struct MembershipPlanner<'a> { + pipeline: &'a mut AppPipeline, + context: &'a mut QueryContext, + pipeline_idx: usize, + input: OutputNodeInfo, + columns: Vec, +} + +impl MembershipPlanner<'_> { + fn rewrite(&mut self, expression: &mut Expr) -> Result<(), PipelineError> { + match expression { + Expr::InSubquery { + expr, + subquery, + negated, + } => { + self.rewrite(expr)?; + validate_query(subquery)?; + let id = self.context.get_next_processor_id(); + let table = format!("__membership_input_{id}"); + let column = format!("__membership_result_{id}"); + let node = format!("membership--{id}"); + // Preserve visible CTEs, but do not expose this query's generated + // table names to a subsequent sibling subquery. + let visible_tables = self.context.pipeline_map.clone(); + let visible_requirements = self.context.replay_requirements.clone(); + let result = query_to_pipeline( + TableInfo { + name: NameOrAlias(table.clone(), None), + override_name: None, + }, + *subquery.clone(), + self.pipeline, + self.context, + self.pipeline_idx, + false, + ); + if let Some(requirements) = self + .context + .replay_requirements + .get(&(self.pipeline_idx, table.clone())) + { + for requirement in requirements { + requirement.require(); + } + } + let values = self + .context + .pipeline_map + .get(&(self.pipeline_idx, table)) + .cloned(); + self.context.pipeline_map = visible_tables; + self.context.replay_requirements = visible_requirements; + result?; + let values = values.ok_or_else(|| { + PipelineError::InvalidQuery("IN subquery has no output relation".into()) + })?; + self.pipeline.add_processor( + Box::new(MembershipProcessorFactory { + id: node.clone(), + operand: *expr.clone(), + negated: *negated, + column: column.clone(), + udfs: self.context.udfs.clone(), + runtime: self.context.runtime.clone(), + }), + node.clone(), + ); + self.pipeline.connect_nodes( + self.input.node.clone(), + self.input.port, + node.clone(), + OUTER_PORT, + ); + self.pipeline + .connect_nodes(values.node, values.port, node.clone(), VALUES_PORT); + self.input = OutputNodeInfo { + node, + port: DEFAULT_PORT_HANDLE, + }; + self.columns.push(column.clone()); + *expression = Expr::Identifier(Ident::new(column)); + } + Expr::BinaryOp { left, right, .. } => { + self.rewrite(left)?; + self.rewrite(right)?; + } + Expr::Nested(expr) + | Expr::UnaryOp { expr, .. } + | Expr::IsNull(expr) + | Expr::IsNotNull(expr) + | Expr::Cast { expr, .. } + | Expr::Extract { expr, .. } => self.rewrite(expr)?, + Expr::InList { expr, list, .. } => { + self.rewrite(expr)?; + for item in list { + self.rewrite(item)?; + } + } + Expr::Like { expr, pattern, .. } => { + self.rewrite(expr)?; + self.rewrite(pattern)?; + } + Expr::Function(function) => { + for argument in &mut function.args { + let (FunctionArg::Named { arg, .. } | FunctionArg::Unnamed(arg)) = argument; + if let FunctionArgExpr::Expr(expr) = arg { + self.rewrite(expr)?; + } + } + } + Expr::Case { + operand, + conditions, + results, + else_result, + } => { + if let Some(operand) = operand { + self.rewrite(operand)?; + } + for item in conditions.iter_mut().chain(results.iter_mut()) { + self.rewrite(item)?; + } + if let Some(otherwise) = else_result { + self.rewrite(otherwise)?; + } + } + // Unsupported expression forms are still rejected by the normal + // ExpressionBuilder rather than being silently discarded here. + _ => {} + } + Ok(()) + } +} + +fn validate_query(query: &Query) -> Result<(), PipelineError> { + validate_query_scope(query, false) +} + +pub(super) fn validate_query_scope(query: &Query, allow_into: bool) -> Result<(), PipelineError> { + if query.fetch.is_some() || !query.locks.is_empty() { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support FETCH or locking clauses".into(), + )); + } + if let Some(with) = &query.with { + for cte in &with.cte_tables { + if !cte.alias.columns.is_empty() { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support CTE column aliases".into(), + )); + } + validate_query(&cte.query)?; + } + } + validate_body(&query.body, allow_into) +} + +fn validate_body(body: &SetExpr, allow_into: bool) -> Result<(), PipelineError> { + let select = match body { + SetExpr::Select(select) => select, + SetExpr::Query(query) => return validate_query_scope(query, allow_into), + SetExpr::SetOperation { + op: SetOperator::Union, + left, + right, + .. + } => { + validate_body(left, false)?; + return validate_body(right, false); + } + _ => { + return Err(PipelineError::InvalidQuery( + "IN subqueries require a SELECT or UNION body".into(), + )) + } + }; + if (!allow_into && select.into.is_some()) || select.distinct.is_some() { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support INTO or SELECT DISTINCT".into(), + )); + } + if select.top.is_some() + || select.qualify.is_some() + || !select.lateral_views.is_empty() + || !select.cluster_by.is_empty() + || !select.distribute_by.is_empty() + || !select.sort_by.is_empty() + || !select.named_window.is_empty() + { + return Err(PipelineError::InvalidQuery( + "IN subquery contains unsupported SELECT modifiers".into(), + )); + } + if select.from.len() != 1 { + return Err(PipelineError::InvalidQuery( + "IN subqueries require one FROM relation, optionally with JOINs".into(), + )); + } + validate_relation(&select.from[0].relation)?; + for join in &select.from[0].joins { + validate_relation(&join.relation)?; + } + for projection in &select.projection { + match projection { + SelectItem::UnnamedExpr(expression) + | SelectItem::ExprWithAlias { + expr: expression, .. + } => { + validate_expression(expression)?; + } + SelectItem::Wildcard(options) | SelectItem::QualifiedWildcard(_, options) + if options.opt_exclude.is_some() + || options.opt_except.is_some() + || options.opt_rename.is_some() + || options.opt_replace.is_some() => + { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support wildcard modifiers".into(), + )); + } + _ => {} + } + } + for expression in select + .selection + .iter() + .chain(select.group_by.iter()) + .chain(select.having.iter()) + { + validate_expression(expression)?; + } + Ok(()) +} + +fn validate_expression(expression: &Expr) -> Result<(), PipelineError> { + match expression { + Expr::InSubquery { expr, subquery, .. } => { + validate_expression(expr)?; + validate_query(subquery)?; + } + Expr::BinaryOp { left, right, .. } => { + validate_expression(left)?; + validate_expression(right)?; + } + Expr::Nested(expr) + | Expr::UnaryOp { expr, .. } + | Expr::IsNull(expr) + | Expr::IsNotNull(expr) + | Expr::Cast { expr, .. } + | Expr::Extract { expr, .. } => { + validate_expression(expr)?; + } + Expr::Interval(interval) => validate_expression(&interval.value)?, + Expr::InList { expr, list, .. } => { + validate_expression(expr)?; + for item in list { + validate_expression(item)?; + } + } + Expr::Like { expr, pattern, .. } => { + validate_expression(expr)?; + validate_expression(pattern)?; + } + Expr::Function(function) => { + if function.distinct || function.over.is_some() || !function.order_by.is_empty() { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support DISTINCT, OVER or ORDER BY function modifiers" + .into(), + )); + } + for argument in &function.args { + let (FunctionArg::Named { arg, .. } | FunctionArg::Unnamed(arg)) = argument; + if let FunctionArgExpr::Expr(expression) = arg { + validate_expression(expression)?; + } + } + } + Expr::Case { + operand, + conditions, + results, + else_result, + } => { + for expression in operand.iter().chain(else_result.iter()) { + validate_expression(expression)?; + } + for expression in conditions.iter().chain(results.iter()) { + validate_expression(expression)?; + } + } + _ => {} + } + Ok(()) +} + +fn validate_relation(relation: &TableFactor) -> Result<(), PipelineError> { + match relation { + TableFactor::Table { + args, + with_hints, + alias, + .. + } => { + if args.is_some() + || !with_hints.is_empty() + || alias + .as_ref() + .is_some_and(|alias| !alias.columns.is_empty()) + { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support table arguments, hints or column aliases".into(), + )); + } + Ok(()) + } + TableFactor::Derived { + lateral, + subquery, + alias, + } => { + if *lateral + || alias + .as_ref() + .is_some_and(|alias| !alias.columns.is_empty()) + { + return Err(PipelineError::InvalidQuery( + "IN subqueries do not support lateral relations or column aliases".into(), + )); + } + validate_query(subquery) + } + _ => Err(PipelineError::InvalidQuery( + "IN subquery contains an unsupported input relation".into(), + )), + } +} diff --git a/dozer-sql/src/builder/mod.rs b/dozer-sql/src/builder/mod.rs index 3ccf3ed356..b5387b2a9a 100644 --- a/dozer-sql/src/builder/mod.rs +++ b/dozer-sql/src/builder/mod.rs @@ -1,7 +1,9 @@ use crate::aggregation::factory::AggregationProcessorFactory; use crate::builder::PipelineError::InvalidQuery; use crate::errors::PipelineError; +use crate::projection::trim::TrimMembershipColumns; use crate::selection::factory::SelectionProcessorFactory; +use crate::selection::replay::ReplayRequirement; use dozer_core::app::AppPipeline; use dozer_core::node::PortHandle; use dozer_core::DEFAULT_PORT_HANDLE; @@ -37,6 +39,7 @@ pub struct OutputNodeInfo { pub struct QueryContext { // Internal tables map, used to store the tables that are created by the queries pipeline_map: HashMap<(usize, String), OutputNodeInfo>, + replay_requirements: HashMap<(usize, String), Vec>, // Output tables map that are marked with "INTO" used to store the tables, these can be exposed to sinks. pub output_tables_map: HashMap, @@ -66,6 +69,7 @@ impl QueryContext { pub fn new(udfs: Vec, runtime: Arc) -> Self { QueryContext { pipeline_map: Default::default(), + replay_requirements: Default::default(), output_tables_map: Default::default(), used_sources: Default::default(), processors_list: Default::default(), @@ -128,6 +132,54 @@ fn query_to_pipeline( query_ctx: &mut QueryContext, pipeline_idx: usize, is_top_select: bool, +) -> Result<(), PipelineError> { + // Each query owns a lexical WITH scope. Keep graph-wide processor IDs and + // shared replay requirements while exposing only this query's output. + let restriction = membership::validate_query_scope(&query, is_top_select).err(); + let visible_tables = query_ctx.pipeline_map.clone(); + let visible_requirements = query_ctx.replay_requirements.clone(); + let output_key = (pipeline_idx, table_info.name.0.clone()); + let result = query_body_to_pipeline( + table_info, + query, + pipeline, + query_ctx, + pipeline_idx, + is_top_select, + ); + let output = query_ctx.pipeline_map.get(&output_key).cloned(); + let requirements = query_ctx.replay_requirements.get(&output_key).cloned(); + if let (Some(error), Some(requirement)) = ( + restriction, + requirements + .as_ref() + .and_then(|requirements| requirements.last()), + ) { + // An outer CTE may have been planned before its membership consumer. + // Preserve unsupported syntax until that consumer requires replay. + requirement.reject_when_required(error.to_string()); + } + query_ctx.pipeline_map = visible_tables; + query_ctx.replay_requirements = visible_requirements; + result?; + if let Some(output) = output { + query_ctx.pipeline_map.insert(output_key.clone(), output); + } + if let Some(requirements) = requirements { + query_ctx + .replay_requirements + .insert(output_key, requirements); + } + Ok(()) +} + +fn query_body_to_pipeline( + table_info: TableInfo, + query: Query, + pipeline: &mut AppPipeline, + query_ctx: &mut QueryContext, + pipeline_idx: usize, + is_top_select: bool, ) -> Result<(), PipelineError> { // return error if there is unsupported syntax if !query.order_by.is_empty() { @@ -150,17 +202,15 @@ fn query_to_pipeline( )); } + let mut local_names = HashSet::new(); for table in with.cte_tables { if table.from.is_some() { return Err(PipelineError::UnsupportedSqlError( UnsupportedSqlError::CteFromError, )); } - let table_name = table.alias.name.to_string(); - if query_ctx - .pipeline_map - .contains_key(&(pipeline_idx, table_name.clone())) - { + let table_name = ExpressionBuilder::normalize_ident(&table.alias.name); + if !local_names.insert(table_name.clone()) { return Err(InvalidQuery(format!( "WITH query name {table_name:?} specified more than once" ))); @@ -179,7 +229,25 @@ fn query_to_pipeline( } }; - match *query.body { + set_body_to_pipeline( + table_info, + *query.body, + pipeline, + query_ctx, + pipeline_idx, + is_top_select, + ) +} + +fn set_body_to_pipeline( + table_info: TableInfo, + body: SetExpr, + pipeline: &mut AppPipeline, + query_ctx: &mut QueryContext, + pipeline_idx: usize, + is_top_select: bool, +) -> Result<(), PipelineError> { + match body { SetExpr::Select(select) => { select_to_pipeline( table_info, @@ -190,47 +258,40 @@ fn query_to_pipeline( is_top_select, )?; } - SetExpr::Query(query) => { - let query_name = format!("subquery_{}", query_ctx.get_next_processor_id()); - let mut ctx = QueryContext::new(query_ctx.udfs.clone(), query_ctx.runtime.clone()); - query_to_pipeline( - TableInfo { - name: NameOrAlias(query_name, None), - override_name: None, - }, - *query, - pipeline, - &mut ctx, - pipeline_idx, - false, //Inside a subquery, so not top select - )? - } + SetExpr::Query(query) => query_to_pipeline( + table_info, + *query, + pipeline, + query_ctx, + pipeline_idx, + is_top_select, + )?, SetExpr::SetOperation { - op, + op: SetOperator::Union, set_quantifier, left, right, - } => match op { - SetOperator::Union => { - set_to_pipeline( - table_info, - left, - right, - set_quantifier, - pipeline, - query_ctx, - pipeline_idx, - is_top_select, - )?; - } - _ => return Err(PipelineError::InvalidOperator(op.to_string())), - }, + } => { + set_to_pipeline( + table_info, + left, + right, + set_quantifier, + pipeline, + query_ctx, + pipeline_idx, + is_top_select, + )?; + } + SetExpr::SetOperation { op, .. } => { + return Err(PipelineError::InvalidOperator(op.to_string())); + } _ => { return Err(PipelineError::UnsupportedSqlError( UnsupportedSqlError::GenericError("Unsupported query body structure".to_string()), )) } - }; + } Ok(()) } @@ -259,7 +320,14 @@ fn select_to_pipeline( let gen_selection_name = format!("select--{}", query_ctx.get_next_processor_id()); let (gen_product_name, product_output_port) = output_node; + let mut replay_requirements = Vec::new(); for (source_name, processor_name, processor_port) in input_nodes { + if let Some(requirements) = query_ctx + .replay_requirements + .get(&(pipeline_idx, source_name.clone())) + { + replay_requirements.extend(requirements.iter().cloned()); + } if let Some(table_info) = query_ctx .pipeline_map .get(&(pipeline_idx, source_name.clone())) @@ -290,28 +358,67 @@ fn select_to_pipeline( query_ctx.runtime.clone(), ); + replay_requirements.push(aggregation.replay_requirement()); pipeline.add_processor(Box::new(aggregation), gen_agg_name.clone()); // Where clause - if let Some(selection) = select.selection { + if let Some(mut selection) = select.selection { + let (input, membership_columns) = membership::plan( + &mut selection, + OutputNodeInfo { + node: gen_product_name, + port: product_output_port, + }, + pipeline, + query_ctx, + pipeline_idx, + )?; let selection = SelectionProcessorFactory::new( gen_selection_name.clone(), selection, query_ctx.udfs.clone(), query_ctx.runtime.clone(), ); + replay_requirements.push(selection.replay_requirement()); + if !membership_columns.is_empty() { + for requirement in &replay_requirements { + requirement.require(); + } + } pipeline.add_processor(Box::new(selection), gen_selection_name.clone()); pipeline.connect_nodes( - gen_product_name, - product_output_port, + input.node, + input.port, gen_selection_name.clone(), DEFAULT_PORT_HANDLE, ); + let selection_output = if membership_columns.is_empty() { + gen_selection_name + } else { + let name = format!( + "membership_projection--{}", + query_ctx.get_next_processor_id() + ); + pipeline.add_processor( + Box::new(TrimMembershipColumns { + id: name.clone(), + columns: membership_columns, + }), + name.clone(), + ); + pipeline.connect_nodes( + gen_selection_name, + DEFAULT_PORT_HANDLE, + name.clone(), + DEFAULT_PORT_HANDLE, + ); + name + }; pipeline.connect_nodes( - gen_selection_name, + selection_output, DEFAULT_PORT_HANDLE, gen_agg_name.clone(), DEFAULT_PORT_HANDLE, @@ -325,6 +432,10 @@ fn select_to_pipeline( ); } + query_ctx.replay_requirements.insert( + (pipeline_idx, table_info.name.0.clone()), + replay_requirements, + ); query_ctx.pipeline_map.insert( (pipeline_idx, table_info.name.0.to_string()), OutputNodeInfo { @@ -382,70 +493,36 @@ fn set_to_pipeline( override_name: None, }; - let _left_pipeline_name = match *left_select { - SetExpr::Select(select) => select_to_pipeline( - left_table_info, - *select, - pipeline, - query_ctx, - pipeline_idx, - is_top_select, - )?, - SetExpr::SetOperation { - op: _, - set_quantifier, - left, - right, - } => set_to_pipeline( - left_table_info, - left, - right, - set_quantifier, - pipeline, - query_ctx, - pipeline_idx, - is_top_select, - )?, - _ => { - return Err(PipelineError::InvalidQuery( - "Invalid UNION left Query".to_string(), - )) - } - }; - - let _right_pipeline_name = match *right_select { - SetExpr::Select(select) => select_to_pipeline( - right_table_info, - *select, - pipeline, - query_ctx, - pipeline_idx, - is_top_select, - )?, - SetExpr::SetOperation { - op: _, - set_quantifier, - left, - right, - } => set_to_pipeline( - right_table_info, - left, - right, - set_quantifier, - pipeline, - query_ctx, - pipeline_idx, - is_top_select, - )?, - _ => { - return Err(PipelineError::InvalidQuery( - "Invalid UNION right Query".to_string(), - )) - } - }; + // UNION operands are internal relations; only the complete set expression + // owns the output name. Parenthesized operands retain their own WITH scope. + set_body_to_pipeline( + left_table_info, + *left_select, + pipeline, + query_ctx, + pipeline_idx, + false, + )?; + set_body_to_pipeline( + right_table_info, + *right_select, + pipeline, + query_ctx, + pipeline_idx, + false, + )?; let mut gen_set_name = format!("set_{}", query_ctx.get_next_processor_id()); + let mut replay_requirements = Vec::new(); + for name in [&gen_left_set_name, &gen_right_set_name] { + if let Some(requirements) = query_ctx + .replay_requirements + .get(&(pipeline_idx, name.clone())) + { + replay_requirements.extend(requirements.iter().cloned()); + } + } let left_pipeline_output_node = query_ctx .pipeline_map .get(&(pipeline_idx, gen_left_set_name)) @@ -460,15 +537,26 @@ fn set_to_pipeline( gen_set_name = table_info.override_name.to_owned().unwrap(); } - let set_proc_fac = SetProcessorFactory::new( - gen_set_name.clone(), - set_quantifier, - pipeline - .flags() - .enable_probabilistic_optimizations - .in_sets - .unwrap_or(false), - ); + let probabilistic = pipeline + .flags() + .enable_probabilistic_optimizations + .in_sets + .unwrap_or(false); + if probabilistic + && matches!( + set_quantifier, + SetQuantifier::None | SetQuantifier::Distinct + ) + { + // Approximate counting cannot retain the original row representation + // for a later retraction. Membership needs exact, replayable inputs. + let requirement = replay_requirements + .last() + .ok_or_else(|| InvalidQuery("UNION input has no replay requirement".into()))?; + requirement.reject_when_required("IN subqueries require exact UNION counting".into()); + } + let set_proc_fac = + SetProcessorFactory::new(gen_set_name.clone(), set_quantifier, probabilistic); pipeline.add_processor(Box::new(set_proc_fac), gen_set_name.clone()); @@ -486,10 +574,10 @@ fn set_to_pipeline( 1, ); - for (_, table_name) in query_ctx.pipeline_map.keys() { - query_ctx.output_tables_map.remove_entry(table_name); - } - + query_ctx.replay_requirements.insert( + (pipeline_idx, table_info.name.0.clone()), + replay_requirements, + ); query_ctx.pipeline_map.insert( (pipeline_idx, table_info.name.0.to_string()), OutputNodeInfo { @@ -498,6 +586,21 @@ fn set_to_pipeline( }, ); + if is_top_select { + let name = table_info + .override_name + .ok_or(PipelineError::MissingIntoClause)?; + if query_ctx.output_tables_map.contains_key(&name) { + return Err(PipelineError::DuplicateIntoClause(name)); + } + query_ctx.output_tables_map.insert( + name, + OutputNodeInfo { + node: gen_set_name.clone(), + port: DEFAULT_PORT_HANDLE, + }, + ); + } Ok(gen_set_name) } @@ -561,6 +664,7 @@ struct ConnectionInfo { mod common; mod from; mod join; +mod membership; mod table_operator; pub use common::string_from_sql_object_name; diff --git a/dozer-sql/src/product/set/operator.rs b/dozer-sql/src/product/set/operator.rs index 1f142f43e3..37cdf271d2 100644 --- a/dozer-sql/src/product/set/operator.rs +++ b/dozer-sql/src/product/set/operator.rs @@ -32,7 +32,7 @@ impl SetOperation { ) -> Result, PipelineError> { match (self.op, self.quantifier) { (SetOperator::Union, SetQuantifier::All) => Ok(vec![(action, record)]), - (SetOperator::Union, SetQuantifier::None) => { + (SetOperator::Union, SetQuantifier::None | SetQuantifier::Distinct) => { self.execute_union(action, record, record_map) } _ => Err(PipelineError::InvalidOperandType(self.op.to_string())), @@ -71,9 +71,14 @@ impl SetOperation { record: Record, record_map: &mut CountingRecordMapEnum, ) -> Result, PipelineError> { - let _count = self.update_map(record.clone(), true, record_map); + // Equality can ignore decimal scale or timestamp offset. Retract the + // representation that was inserted, so later expressions see the same row. + let representative = record_map + .representative(&record) + .unwrap_or_else(|| record.clone()); + let _count = self.update_map(record, true, record_map); if _count == 0 { - Ok(vec![(action, record)]) + Ok(vec![(action, representative)]) } else { Ok(vec![]) } @@ -94,3 +99,44 @@ impl SetOperation { record_map.estimate_count(&record) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::product::set::record_map::AccurateCountingRecordMap; + use dozer_types::{rust_decimal::Decimal, types::Field}; + + #[test] + fn union_retracts_the_first_equal_decimal_representation() { + let first = Record::new(vec![Field::Decimal(Decimal::new(10, 1))]); + let second = Record::new(vec![Field::Decimal(Decimal::new(100, 2))]); + assert_eq!(first, second); + let encoded = + |row: &Record| bincode::encode_to_vec(row, bincode::config::standard()).unwrap(); + assert_ne!(encoded(&first), encoded(&second)); + for quantifier in [SetQuantifier::None, SetQuantifier::Distinct] { + let op = SetOperation { + op: SetOperator::Union, + quantifier, + }; + let mut map: CountingRecordMapEnum = AccurateCountingRecordMap::new().unwrap().into(); + let inserted = op + .execute(SetAction::Insert, first.clone(), &mut map) + .unwrap(); + assert!(op + .execute(SetAction::Insert, second.clone(), &mut map) + .unwrap() + .is_empty()); + assert!(op + .execute(SetAction::Delete, first.clone(), &mut map) + .unwrap() + .is_empty()); + let deleted = op + .execute(SetAction::Delete, second.clone(), &mut map) + .unwrap(); + assert_eq!(deleted.len(), 1); + assert_eq!(deleted[0].0, SetAction::Delete); + assert_eq!(encoded(&deleted[0].1), encoded(&inserted[0].1)); + } + } +} diff --git a/dozer-sql/src/product/set/record_map/mod.rs b/dozer-sql/src/product/set/record_map/mod.rs index 1d4119e18f..7925279812 100644 --- a/dozer-sql/src/product/set/record_map/mod.rs +++ b/dozer-sql/src/product/set/record_map/mod.rs @@ -24,6 +24,9 @@ pub trait CountingRecordMap { /// Depending on the implementation, this number may not be accurate. fn estimate_count(&self, record: &Record) -> u64; + /// Returns the stored representation when exact counting retains it. + fn representative(&self, record: &Record) -> Option; + /// Clears the map, removing all records. fn clear(&mut self); } @@ -62,6 +65,12 @@ impl CountingRecordMap for AccurateCountingRecordMap { self.map.get(record).copied().unwrap_or(0) } + fn representative(&self, record: &Record) -> Option { + self.map + .get_key_value(record) + .map(|(stored, _)| stored.clone()) + } + fn clear(&mut self) { self.map.clear(); } @@ -100,6 +109,10 @@ impl CountingRecordMap for ProbabilisticCountingRecordMap { self.map.estimate_count(record) as u64 } + fn representative(&self, _record: &Record) -> Option { + None + } + fn clear(&mut self) { self.map.clear(); } diff --git a/dozer-sql/src/product/set/set_factory.rs b/dozer-sql/src/product/set/set_factory.rs index 6f4fd442f4..64c3e7e0d7 100644 --- a/dozer-sql/src/product/set/set_factory.rs +++ b/dozer-sql/src/product/set/set_factory.rs @@ -64,7 +64,8 @@ impl ProcessorFactory for SetProcessorFactory { let output_schema = Schema { fields: output_columns, - primary_index: input_schemas[&0].primary_index.clone(), + // A key from either input is not necessarily unique across UNION. + primary_index: vec![], }; Ok(output_schema) @@ -90,27 +91,86 @@ impl ProcessorFactory for SetProcessorFactory { fn validate_set_operation_input_schemas( input_schemas: &HashMap, ) -> Result, PipelineError> { - let mut left_columns = input_schemas[&0].fields.clone(); - let mut right_columns = input_schemas[&0].fields.clone(); - - left_columns.sort(); - right_columns.sort(); + let left = input_schemas + .get(&0) + .ok_or(PipelineError::InvalidPortHandle(0))?; + let right = input_schemas + .get(&1) + .ok_or(PipelineError::InvalidPortHandle(1))?; + if left.fields.len() != right.fields.len() { + return Err(PipelineError::SetError(SetError::InvalidInputSchemas)); + } + left.fields + .iter() + .zip(&right.fields) + .map(|(left, right)| { + if left.typ != right.typ { + return Err(PipelineError::SetError(SetError::InvalidInputSchemas)); + } + // UNION aligns values by position and keeps the left branch's names. + Ok(FieldDefinition::new( + left.name.clone(), + left.typ, + left.nullable || right.nullable, + SourceDefinition::Dynamic, + )) + }) + .collect() +} - let mut output_fields = Vec::new(); - for (left, right) in left_columns.iter().zip(right_columns.iter()) { - if !is_similar_fields(left, right) { - return Err(PipelineError::SetError(SetError::InvalidInputSchemas)); +#[cfg(test)] +mod tests { + use super::*; + use dozer_types::types::FieldType; + + fn schema(fields: &[(&str, FieldType, bool)]) -> Schema { + Schema { + fields: fields + .iter() + .map(|(name, typ, nullable)| { + FieldDefinition::new((*name).into(), *typ, *nullable, SourceDefinition::Dynamic) + }) + .collect(), + primary_index: vec![], } - output_fields.push(FieldDefinition::new( - left.name.clone(), - left.typ, - left.nullable, - SourceDefinition::Dynamic, - )); } - Ok(output_fields) -} -fn is_similar_fields(left: &FieldDefinition, right: &FieldDefinition) -> bool { - left.name == right.name && left.typ == right.typ && left.nullable == right.nullable + #[test] + fn union_schema_preserves_position_and_combines_nullability() { + let left = schema(&[ + ("z", FieldType::Int, false), + ("a", FieldType::String, false), + ]); + let right = schema(&[ + ("other", FieldType::Int, true), + ("name", FieldType::String, false), + ]); + let output = + validate_set_operation_input_schemas(&HashMap::from([(0, left), (1, right)])).unwrap(); + assert_eq!( + output, + schema(&[("z", FieldType::Int, true), ("a", FieldType::String, false)]).fields + ); + } + + #[test] + fn union_schema_rejects_wrong_right_branch_width_or_type() { + let left = schema(&[("value", FieldType::Int, true)]); + for right in [ + schema(&[]), + schema(&[("value", FieldType::String, true)]), + schema(&[ + ("value", FieldType::Int, true), + ("extra", FieldType::Int, true), + ]), + ] { + assert!(matches!( + validate_set_operation_input_schemas(&HashMap::from([ + (0, left.clone()), + (1, right) + ])), + Err(PipelineError::SetError(SetError::InvalidInputSchemas)) + )); + } + } } diff --git a/dozer-sql/src/projection/mod.rs b/dozer-sql/src/projection/mod.rs index 9a12ba12cc..0816cdbf7e 100644 --- a/dozer-sql/src/projection/mod.rs +++ b/dozer-sql/src/projection/mod.rs @@ -1,2 +1,3 @@ pub mod factory; pub mod processor; +pub(crate) mod trim; diff --git a/dozer-sql/src/projection/trim.rs b/dozer-sql/src/projection/trim.rs new file mode 100644 index 0000000000..cadfa32372 --- /dev/null +++ b/dozer-sql/src/projection/trim.rs @@ -0,0 +1,87 @@ +use std::collections::HashMap; + +use dozer_core::{ + event::EventHub, + node::{PortHandle, Processor, ProcessorFactory}, + DEFAULT_PORT_HANDLE, +}; +use dozer_sql_expression::execution::Expression; +use dozer_types::{errors::internal::BoxedError, tonic::async_trait, types::Schema}; + +use crate::errors::PipelineError; + +use super::processor::ProjectionProcessor; + +/// Restore the input schema before wildcard expansion or aggregation. +#[derive(Debug)] +pub(crate) struct TrimMembershipColumns { + pub id: String, + pub columns: Vec, +} + +impl TrimMembershipColumns { + fn width(&self, schema: &Schema) -> Result { + let width = schema.fields.len().checked_sub(self.columns.len()); + match width { + Some(width) + if schema.fields[width..] + .iter() + .map(|field| &field.name) + .eq(self.columns.iter()) => + { + Ok(width) + } + _ => Err(PipelineError::InvalidQuery( + "IN membership columns do not match the planned schema".into(), + )), + } + } +} + +#[async_trait] +impl ProcessorFactory for TrimMembershipColumns { + fn id(&self) -> String { + self.id.clone() + } + fn type_name(&self) -> String { + "MembershipProjection".into() + } + fn get_input_ports(&self) -> Vec { + vec![DEFAULT_PORT_HANDLE] + } + fn get_output_ports(&self) -> Vec { + vec![DEFAULT_PORT_HANDLE] + } + + async fn get_output_schema( + &self, + _port: &PortHandle, + inputs: &HashMap, + ) -> Result { + let mut schema = inputs + .get(&DEFAULT_PORT_HANDLE) + .ok_or(PipelineError::InvalidPortHandle(DEFAULT_PORT_HANDLE))? + .clone(); + let width = self.width(&schema)?; + schema.fields.truncate(width); + Ok(schema) + } + + async fn build( + &self, + inputs: HashMap, + _outputs: HashMap, + _events: EventHub, + ) -> Result, BoxedError> { + let schema = inputs + .get(&DEFAULT_PORT_HANDLE) + .ok_or(PipelineError::InvalidPortHandle(DEFAULT_PORT_HANDLE))?; + let expressions = (0..self.width(schema)?) + .map(|index| Expression::Column { index }) + .collect(); + Ok(Box::new(ProjectionProcessor::new( + schema.clone(), + expressions, + )?)) + } +} diff --git a/dozer-sql/src/selection/factory.rs b/dozer-sql/src/selection/factory.rs index ab5567584c..0045f0f58c 100644 --- a/dozer-sql/src/selection/factory.rs +++ b/dozer-sql/src/selection/factory.rs @@ -13,6 +13,7 @@ use dozer_types::{models::udf_config::UdfConfig, tonic::async_trait}; use tokio::runtime::Runtime; use super::processor::SelectionProcessor; +use super::replay::ReplayRequirement; #[derive(Debug)] pub struct SelectionProcessorFactory { @@ -20,6 +21,7 @@ pub struct SelectionProcessorFactory { id: String, udfs: Vec, runtime: Arc, + replay: ReplayRequirement, } impl SelectionProcessorFactory { @@ -35,8 +37,13 @@ impl SelectionProcessorFactory { id, udfs: udf_config, runtime, + replay: ReplayRequirement::default(), } } + + pub(crate) fn replay_requirement(&self) -> ReplayRequirement { + self.replay.clone() + } } #[async_trait] @@ -80,10 +87,13 @@ impl ProcessorFactory for SelectionProcessorFactory { .build(false, &self.statement, schema, &self.udfs) .await { - Ok(expression) => Ok(Box::new(SelectionProcessor::new( - schema.clone(), - expression, - )?)), + Ok(expression) => { + self.replay.validate(&expression)?; + Ok(Box::new(SelectionProcessor::new( + schema.clone(), + expression, + )?)) + } Err(e) => Err(e.into()), } } diff --git a/dozer-sql/src/selection/membership.rs b/dozer-sql/src/selection/membership.rs new file mode 100644 index 0000000000..6ed837e18c --- /dev/null +++ b/dozer-sql/src/selection/membership.rs @@ -0,0 +1,477 @@ +//! Incremental membership of an outer expression in a single-column relation. +//! +//! Every outer row is retained, including rows whose membership is false or +//! unknown. The extra nullable boolean is consumed by a downstream selection; +//! this keeps the same operator usable inside compound WHERE predicates. + +use std::{collections::HashMap, sync::Arc}; + +use dozer_core::{ + channels::ProcessorChannelForwarder, + epoch::Epoch, + event::EventHub, + node::{PortHandle, Processor, ProcessorFactory}, + DEFAULT_PORT_HANDLE, +}; +use dozer_sql_expression::{ + builder::ExpressionBuilder, execution::Expression, operator::BinaryOperatorType, + sqlparser::ast::Expr, +}; +use dozer_types::{ + errors::internal::BoxedError, + models::udf_config::UdfConfig, + tonic::async_trait, + types::{Field, FieldDefinition, FieldType, Operation, Record, Schema, TableOperation}, +}; +use tokio::runtime::Runtime; + +use crate::errors::PipelineError; + +pub(crate) const OUTER_PORT: PortHandle = 0; +pub(crate) const VALUES_PORT: PortHandle = 1; + +#[derive(Debug)] +pub(crate) struct MembershipProcessorFactory { + pub id: String, + pub operand: Expr, + pub negated: bool, + pub column: String, + pub udfs: Vec, + pub runtime: Arc, +} + +impl MembershipProcessorFactory { + fn outer_schema<'a>( + &self, + inputs: &'a HashMap, + ) -> Result<&'a Schema, PipelineError> { + let outer = inputs + .get(&OUTER_PORT) + .ok_or(PipelineError::InvalidPortHandle(OUTER_PORT))?; + let values = inputs + .get(&VALUES_PORT) + .ok_or(PipelineError::InvalidPortHandle(VALUES_PORT))?; + if values.fields.len() != 1 { + return Err(PipelineError::InvalidQuery( + "An IN subquery must return one column".into(), + )); + } + if outer.fields.iter().any(|field| field.name == self.column) { + return Err(PipelineError::InvalidQuery(format!( + "IN subquery output column already exists: {}", + self.column + ))); + } + Ok(outer) + } +} + +#[async_trait] +impl ProcessorFactory for MembershipProcessorFactory { + fn id(&self) -> String { + self.id.clone() + } + + fn type_name(&self) -> String { + "Membership".into() + } + + fn get_input_ports(&self) -> Vec { + vec![OUTER_PORT, VALUES_PORT] + } + + fn get_output_ports(&self) -> Vec { + vec![DEFAULT_PORT_HANDLE] + } + + async fn get_output_schema( + &self, + _output: &PortHandle, + inputs: &HashMap, + ) -> Result { + let mut schema = self.outer_schema(inputs)?.clone(); + schema.fields.push(FieldDefinition::new( + self.column.clone(), + FieldType::Boolean, + true, + Default::default(), + )); + Ok(schema) + } + + async fn build( + &self, + inputs: HashMap, + _outputs: HashMap, + _events: EventHub, + ) -> Result, BoxedError> { + let schema = self.outer_schema(&inputs)?; + let operand = ExpressionBuilder::new(schema.fields.len(), self.runtime.clone()) + .build(false, &self.operand, schema, &self.udfs) + .await?; + Ok(Box::new(MembershipProcessor::new( + schema.clone(), + operand, + self.negated, + )?)) + } +} + +#[derive(Debug)] +struct RetainedRow { + record: Record, + operand: Field, + copies: usize, + matches: usize, + unknown: usize, +} + +#[derive(Debug)] +pub(crate) struct MembershipProcessor { + schema: Schema, + operand: Expression, + negated: bool, + rows: HashMap, RetainedRow>, + values: HashMap, + value_count: usize, +} + +impl MembershipProcessor { + pub(crate) fn new( + schema: Schema, + operand: Expression, + negated: bool, + ) -> Result { + validate_operand(&operand)?; + Ok(Self { + schema, + operand, + negated, + rows: HashMap::new(), + values: HashMap::new(), + value_count: 0, + }) + } + + fn result(&self, matches: usize, unknown: usize, total: usize) -> Field { + if total == 0 { + Field::Boolean(self.negated) + } else if matches != 0 { + Field::Boolean(!self.negated) + } else if unknown != 0 { + Field::Null + } else { + Field::Boolean(self.negated) + } + } + + fn equal(&self, left: &Field, right: &Field) -> Result, PipelineError> { + if left == &Field::Null || right == &Field::Null { + return Ok(None); + } + // Field equality distinguishes numeric variants; SQL equality performs + // the same conversions as an ordinary WHERE comparison. + let value = BinaryOperatorType::Eq.evaluate( + &self.schema, + &mut Expression::Literal(left.clone()), + &mut Expression::Literal(right.clone()), + &Record::default(), + )?; + match value { + Field::Boolean(value) => Ok(Some(value)), + Field::Null => Ok(None), + _ => Err(PipelineError::InvalidQuery( + "IN operands have no boolean equality comparison".into(), + )), + } + } + + fn prepare_row(&mut self, row: &Record) -> Result<(Vec, RetainedRow), PipelineError> { + validate_lifetime(row)?; + if row.values.len() != self.schema.fields.len() { + return Err(PipelineError::InvalidValue( + "IN outer row does not match its schema".into(), + )); + } + let operand = self.operand.evaluate(row, &self.schema)?; + let mut matches = 0; + let mut unknown = 0; + for (value, copies) in &self.values { + match self.equal(&operand, value)? { + Some(true) => matches += copies, + None => unknown += copies, + Some(false) => {} + } + } + Ok(( + row_key(row)?, + RetainedRow { + record: row.clone(), + operand, + copies: 1, + matches, + unknown, + }, + )) + } + + fn marked(&self, row: &Record, retained: &RetainedRow) -> Record { + append_result( + row, + self.result(retained.matches, retained.unknown, self.value_count), + ) + } + + fn retained(&self, row: &Record) -> Result<&RetainedRow, PipelineError> { + validate_lifetime(row)?; + self.rows.get(&row_key(row)?).ok_or_else(|| { + PipelineError::InvalidValue("IN received a change for an unknown outer row".into()) + }) + } + + fn remove_row(&mut self, key: &[u8]) { + let retained = self.rows.get_mut(key).expect("outer row validated"); + retained.copies -= 1; + if retained.copies == 0 { + self.rows.remove(key); + } + } + + fn add_row(&mut self, key: Vec, retained: RetainedRow) { + if let Some(existing) = self.rows.get_mut(&key) { + existing.copies += 1; + } else { + self.rows.insert(key, retained); + } + } + + fn change_outer(&mut self, operation: Operation) -> Result { + match operation { + Operation::Insert { new } => { + let (key, retained) = self.prepare_row(&new)?; + let marked = self.marked(&new, &retained); + self.add_row(key, retained); + Ok(Operation::Insert { new: marked }) + } + Operation::Delete { old } => { + let key = row_key(&old)?; + let marked = self.marked(&old, self.retained(&old)?); + self.remove_row(&key); + Ok(Operation::Delete { old: marked }) + } + Operation::Update { old, new } => { + let old_key = row_key(&old)?; + let previous = self.marked(&old, self.retained(&old)?); + let (new_key, retained) = self.prepare_row(&new)?; + let next = self.marked(&new, &retained); + self.remove_row(&old_key); + self.add_row(new_key, retained); + Ok(Operation::Update { + old: previous, + new: next, + }) + } + Operation::BatchInsert { new } => { + // Evaluate the whole batch before retaining or emitting any row. + let mut prepared = Vec::with_capacity(new.len()); + for row in new { + prepared.push(self.prepare_row(&row)?); + } + let mut output = Vec::with_capacity(prepared.len()); + for (key, retained) in prepared { + output.push(self.marked(&retained.record, &retained)); + self.add_row(key, retained); + } + Ok(Operation::BatchInsert { new: output }) + } + } + } + + fn change_values(&mut self, operation: Operation) -> Result, PipelineError> { + let (removed, added) = match operation { + Operation::Insert { new } => (Vec::new(), vec![new]), + Operation::Delete { old } => (vec![old], Vec::new()), + Operation::Update { old, new } => (vec![old], vec![new]), + Operation::BatchInsert { new } => (Vec::new(), new), + }; + let total = self + .value_count + .checked_sub(removed.len()) + .and_then(|count| count.checked_add(added.len())) + .ok_or_else(|| PipelineError::InvalidValue("Invalid IN inner row count".into()))?; + let mut replacements = HashMap::new(); + for row in &removed { + let value = inner_value(row)?; + let count = replacements + .entry(value.clone()) + .or_insert_with(|| self.values.get(value).copied().unwrap_or(0)); + *count = count.checked_sub(1).ok_or_else(|| { + PipelineError::InvalidValue("IN received a delete for an unknown inner row".into()) + })?; + } + for row in &added { + let value = inner_value(row)?; + *replacements + .entry(value.clone()) + .or_insert_with(|| self.values.get(value).copied().unwrap_or(0)) += 1; + } + let mut changes = Vec::with_capacity(self.rows.len()); + // Calculate every comparison before mutating either input's state. An + // incompatible value must not partially apply an UPDATE or batch. + for (key, retained) in &self.rows { + let mut matches = retained.matches; + let mut unknown = retained.unknown; + for (value, next_count) in &replacements { + let old_count = self.values.get(value).copied().unwrap_or(0); + if old_count != *next_count { + match self.equal(&retained.operand, value)? { + Some(true) => matches = matches - old_count + next_count, + None => unknown = unknown - old_count + next_count, + Some(false) => {} + } + } + } + let before = self.result(retained.matches, retained.unknown, self.value_count); + let after = self.result(matches, unknown, total); + changes.push((key.clone(), matches, unknown, before, after)); + } + let mut output = Vec::new(); + for (key, matches, unknown, before, after) in changes { + let retained = self.rows.get_mut(&key).expect("retained row"); + retained.matches = matches; + retained.unknown = unknown; + if before != after { + let update = Operation::Update { + old: append_result(&retained.record, before), + new: append_result(&retained.record, after), + }; + output.extend(std::iter::repeat(update).take(retained.copies)); + } + } + for (value, count) in replacements { + if count == 0 { + self.values.remove(&value); + } else { + self.values.insert(value, count); + } + } + self.value_count = total; + Ok(output) + } +} + +fn row_key(row: &Record) -> Result, PipelineError> { + // Record equality ignores decimal scale and timestamp offsets, whereas + // scalar expressions such as CAST(... AS STRING) can distinguish them. + bincode::encode_to_vec(row, bincode::config::standard()) + .map_err(|error| PipelineError::InternalError(Box::new(error))) +} + +pub(crate) fn validate_operand(expression: &Expression) -> Result<(), PipelineError> { + // The retained operand must be reproducible for identical outer rows. + // Unknown extension functions are excluded until they declare determinism. + use Expression::*; + match expression { + Column { .. } | Literal(_) => Ok(()), + UnaryOperator { arg, .. } + | DateTimeFunction { arg, .. } + | Cast { arg, .. } + | IsNull { arg } + | IsNotNull { arg } => validate_operand(arg), + BinaryOperator { left, right, .. } => { + validate_operand(left)?; + validate_operand(right) + } + ScalarFunction { args, .. } + | GeoFunction { args, .. } + | ConditionalExpression { args, .. } + | Json { args, .. } => args.iter().try_for_each(validate_operand), + Trim { arg, what, .. } => { + validate_operand(arg)?; + what.iter() + .try_for_each(|expression| validate_operand(expression)) + } + Like { arg, pattern, .. } => { + validate_operand(arg)?; + validate_operand(pattern) + } + InList { expr, list, .. } => { + validate_operand(expr)?; + list.iter().try_for_each(validate_operand) + } + Case { + operand, + conditions, + results, + else_result, + } => { + operand + .iter() + .try_for_each(|expression| validate_operand(expression))?; + conditions.iter().try_for_each(validate_operand)?; + results.iter().try_for_each(validate_operand)?; + else_result + .iter() + .try_for_each(|expression| validate_operand(expression)) + } + _ => Err(PipelineError::InvalidQuery( + "IN subquery operands must be deterministic scalar expressions".into(), + )), + } +} + +fn append_result(row: &Record, result: Field) -> Record { + let mut output = row.clone(); + output.values.push(result); + output +} + +fn validate_lifetime(row: &Record) -> Result<(), PipelineError> { + if row.lifetime.is_some() { + // The current TTL path expires join state without sending downstream + // deletions. Membership cannot stay correct under that contract. + return Err(PipelineError::InvalidQuery( + "IN subqueries require explicit deletion events; TTL inputs are unsupported".into(), + )); + } + Ok(()) +} + +fn inner_value(row: &Record) -> Result<&Field, PipelineError> { + validate_lifetime(row)?; + if row.values.len() != 1 { + return Err(PipelineError::InvalidValue( + "IN inner row must contain one value".into(), + )); + } + Ok(&row.values[0]) +} + +impl Processor for MembershipProcessor { + fn commit(&self, _epoch: &Epoch) -> Result<(), BoxedError> { + Ok(()) + } + + fn process( + &mut self, + operation: TableOperation, + forwarder: &mut dyn ProcessorChannelForwarder, + ) -> Result<(), BoxedError> { + let output = match operation.port { + OUTER_PORT => vec![self.change_outer(operation.op)?], + VALUES_PORT => self.change_values(operation.op)?, + port => return Err(PipelineError::InvalidPortHandle(port).into()), + }; + for op in output { + forwarder.send(TableOperation { + id: operation.id, + port: DEFAULT_PORT_HANDLE, + op, + }); + } + Ok(()) + } +} + +#[cfg(test)] +mod tests; diff --git a/dozer-sql/src/selection/membership/tests.rs b/dozer-sql/src/selection/membership/tests.rs new file mode 100644 index 0000000000..14312629cd --- /dev/null +++ b/dozer-sql/src/selection/membership/tests.rs @@ -0,0 +1,705 @@ +use super::*; +use crate::selection::processor::SelectionProcessor; +use dozer_types::{node::OpIdentifier, types::Lifetime}; + +#[derive(Default)] +struct Output(Vec); + +impl ProcessorChannelForwarder for Output { + fn send(&mut self, operation: TableOperation) { + self.0.push(operation); + } +} + +fn schema() -> Schema { + Schema { + fields: vec![ + FieldDefinition::new("id".into(), FieldType::Int, false, Default::default()), + FieldDefinition::new("value".into(), FieldType::Int, true, Default::default()), + ], + primary_index: vec![0], + } +} + +fn processor(negated: bool) -> MembershipProcessor { + MembershipProcessor::new(schema(), Expression::Column { index: 1 }, negated).unwrap() +} + +fn outer(id: i64, value: Field) -> Record { + Record::new(vec![Field::Int(id), value]) +} + +fn inner(value: Field) -> Record { + Record::new(vec![value]) +} + +fn insert(row: Record) -> Operation { + Operation::Insert { new: row } +} + +fn delete(row: Record) -> Operation { + Operation::Delete { old: row } +} + +fn feed( + processor: &mut MembershipProcessor, + port: PortHandle, + operation: Operation, +) -> Result, BoxedError> { + let id = Some(OpIdentifier::new(19, 7)); + let mut output = Output::default(); + let result = processor.process( + TableOperation { + id, + port, + op: operation, + }, + &mut output, + ); + if let Err(error) = result { + assert!( + output.0.is_empty(), + "a failed operation emitted partial output" + ); + return Err(error); + } + for operation in &output.0 { + assert_eq!(operation.port, DEFAULT_PORT_HANDLE); + assert_eq!(operation.id, id); + } + Ok(output.0.into_iter().map(|operation| operation.op).collect()) +} + +fn marked(row: &Record, mark: Field) -> Record { + let mut row = row.clone(); + row.values.push(mark); + row +} + +fn transition(row: &Record, before: Field, after: Field) -> Operation { + Operation::Update { + old: marked(row, before), + new: marked(row, after), + } +} + +fn selected(operations: Vec) -> Vec { + let mut schema = schema(); + schema.fields.push(FieldDefinition::new( + "member".into(), + FieldType::Boolean, + true, + Default::default(), + )); + let mut selection = SelectionProcessor::new(schema, Expression::Column { index: 2 }).unwrap(); + let mut output = Output::default(); + for op in operations { + selection + .process( + TableOperation { + id: None, + port: DEFAULT_PORT_HANDLE, + op, + }, + &mut output, + ) + .unwrap(); + } + output.0.into_iter().map(|operation| operation.op).collect() +} + +#[test] +fn membership_truth_table_including_null_operand_and_empty_relation() { + let cases = [ + (vec![], Field::Boolean(false), Field::Boolean(false)), + (vec![Field::Int(7)], Field::Boolean(true), Field::Null), + (vec![Field::Int(9)], Field::Boolean(false), Field::Null), + (vec![Field::Null], Field::Null, Field::Null), + ( + vec![Field::Int(7), Field::Null], + Field::Boolean(true), + Field::Null, + ), + (vec![Field::Int(9), Field::Null], Field::Null, Field::Null), + ]; + for negated in [false, true] { + for (values, expected, expected_null) in &cases { + let mut p = processor(negated); + feed( + &mut p, + VALUES_PORT, + Operation::BatchInsert { + new: values.iter().cloned().map(inner).collect(), + }, + ) + .unwrap(); + for (value, expected) in [(Field::Int(7), expected), (Field::Null, expected_null)] { + let expected = match expected { + Field::Boolean(value) => Field::Boolean(value ^ negated), + _ => Field::Null, + }; + let row = outer(1, value); + assert_eq!( + feed(&mut p, OUTER_PORT, insert(row.clone())).unwrap(), + vec![insert(marked(&row, expected))], + "negated={negated}, values={values:?}" + ); + } + } + } +} + +#[test] +fn duplicate_inner_rows_do_not_multiply_outer_rows() { + let mut p = processor(false); + let row = outer(1, Field::Int(7)); + feed( + &mut p, + OUTER_PORT, + Operation::BatchInsert { + new: vec![row.clone(), row.clone()], + }, + ) + .unwrap(); + let changes = feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))).unwrap(); + assert_eq!( + changes, + vec![transition(&row, Field::Boolean(false), Field::Boolean(true)); 2] + ); + assert!(feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))) + .unwrap() + .is_empty()); + assert!(feed(&mut p, VALUES_PORT, delete(inner(Field::Int(7)))) + .unwrap() + .is_empty()); + feed(&mut p, OUTER_PORT, delete(row.clone())).unwrap(); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(Field::Int(7)))).unwrap(), + vec![transition( + &row, + Field::Boolean(true), + Field::Boolean(false) + )] + ); + feed(&mut p, OUTER_PORT, delete(row)).unwrap(); + assert!(p.rows.is_empty()); + assert!(p.values.is_empty()); +} + +#[test] +fn outer_updates_produce_all_four_selection_transitions() { + let mut p = processor(false); + feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))).unwrap(); + let mut old = outer(1, Field::Int(9)); + assert!(selected(feed(&mut p, OUTER_PORT, insert(old.clone())).unwrap()).is_empty()); + for (id, value, expected_kind) in [(2, 8, 0), (3, 7, 1), (4, 7, 2), (5, 8, 3)] { + let new = outer(id, Field::Int(value)); + let output = selected( + feed( + &mut p, + OUTER_PORT, + Operation::Update { + old: old.clone(), + new: new.clone(), + }, + ) + .unwrap(), + ); + let expected = match expected_kind { + 0 => vec![], + 1 => vec![insert(marked(&new, Field::Boolean(true)))], + 2 => vec![Operation::Update { + old: marked(&old, Field::Boolean(true)), + new: marked(&new, Field::Boolean(true)), + }], + 3 => vec![delete(marked(&old, Field::Boolean(true)))], + _ => unreachable!(), + }; + assert_eq!(output, expected); + old = new; + } +} + +#[test] +fn inner_update_and_batch_have_no_intermediate_membership_event() { + let mut p = processor(true); + let row = outer(1, Field::Int(7)); + feed(&mut p, OUTER_PORT, insert(row.clone())).unwrap(); + let batch = feed( + &mut p, + VALUES_PORT, + Operation::BatchInsert { + new: vec![inner(Field::Null), inner(Field::Int(7))], + }, + ) + .unwrap(); + assert_eq!( + batch, + vec![transition( + &row, + Field::Boolean(true), + Field::Boolean(false) + )] + ); + assert!(feed( + &mut p, + VALUES_PORT, + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(7)), + }, + ) + .unwrap() + .is_empty()); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(Field::Int(7)))).unwrap(), + vec![transition(&row, Field::Boolean(false), Field::Null)] + ); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(Field::Null))).unwrap(), + vec![transition(&row, Field::Null, Field::Boolean(true))] + ); +} + +#[test] +fn membership_uses_sql_numeric_and_text_equality() { + for (left, right) in [ + (Field::Int(7), Field::Float(7.0.into())), + (Field::UInt(7), Field::Int(7)), + (Field::String("seven".into()), Field::Text("seven".into())), + ] { + let mut p = processor(false); + let row = outer(1, left); + feed(&mut p, OUTER_PORT, insert(row.clone())).unwrap(); + assert_eq!( + feed(&mut p, VALUES_PORT, insert(inner(right.clone()))).unwrap(), + vec![transition( + &row, + Field::Boolean(false), + Field::Boolean(true) + )] + ); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(right))).unwrap(), + vec![transition( + &row, + Field::Boolean(true), + Field::Boolean(false) + )] + ); + } +} + +#[test] +fn replacing_equal_values_of_different_types_is_atomic() { + let mut p = processor(false); + let row = outer(1, Field::Int(7)); + feed(&mut p, OUTER_PORT, insert(row)).unwrap(); + feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))).unwrap(); + assert!(feed( + &mut p, + VALUES_PORT, + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Float(7.0.into())), + }, + ) + .unwrap() + .is_empty()); + assert_eq!(p.value_count, 1); +} + +#[test] +fn malformed_changes_leave_state_and_output_untouched() { + let mut p = processor(false); + let row = outer(1, Field::Int(7)); + feed(&mut p, OUTER_PORT, insert(row.clone())).unwrap(); + feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))).unwrap(); + let bad = inner(Field::Binary(vec![7])); + for operation in [ + Operation::Update { + old: inner(Field::Int(7)), + new: bad.clone(), + }, + Operation::BatchInsert { + new: vec![inner(Field::Int(9)), bad], + }, + delete(inner(Field::Int(9))), + Operation::Update { + old: inner(Field::Int(9)), + new: inner(Field::Int(9)), + }, + insert(Record::new(vec![Field::Int(1), Field::Int(2)])), + ] { + assert!(feed(&mut p, VALUES_PORT, operation).is_err()); + assert_eq!(p.values, HashMap::from([(Field::Int(7), 1)])); + assert_eq!(p.value_count, 1); + assert_eq!(p.retained(&row).unwrap().matches, 1); + } + assert!(feed( + &mut p, + OUTER_PORT, + Operation::BatchInsert { + new: vec![outer(2, Field::Int(7)), Record::new(vec![])], + }, + ) + .is_err()); + assert_eq!(p.rows.len(), 1); + assert!(feed( + &mut p, + OUTER_PORT, + Operation::Update { + old: row.clone(), + new: outer(2, Field::Binary(vec![7])), + }, + ) + .is_err()); + assert_eq!(p.retained(&row).unwrap().copies, 1); + assert!(feed(&mut p, OUTER_PORT, delete(outer(99, Field::Int(7)))).is_err()); + assert!(feed(&mut p, 9, insert(row)).is_err()); +} + +#[test] +fn ttl_rejection_does_not_partially_apply_a_batch() { + let mut p = processor(false); + let mut expired = inner(Field::Int(7)); + expired.lifetime = Some(Lifetime { + reference: dozer_types::chrono::Utc::now().into(), + duration: std::time::Duration::from_secs(10), + }); + assert!(feed( + &mut p, + VALUES_PORT, + Operation::BatchInsert { + new: vec![inner(Field::Int(1)), expired], + }, + ) + .is_err()); + assert!(p.values.is_empty()); + assert_eq!(p.value_count, 0); +} + +#[test] +fn unknown_comparisons_are_counted_even_without_null_values() { + let mut p = processor(false); + let row = outer(1, Field::Boolean(true)); + feed(&mut p, OUTER_PORT, insert(row.clone())).unwrap(); + assert_eq!( + feed(&mut p, VALUES_PORT, insert(inner(Field::Binary(vec![1])))).unwrap(), + vec![transition(&row, Field::Boolean(false), Field::Null)] + ); + assert_eq!( + feed(&mut p, VALUES_PORT, insert(inner(Field::Boolean(true)))).unwrap(), + vec![transition(&row, Field::Null, Field::Boolean(true))] + ); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(Field::Boolean(true)))).unwrap(), + vec![transition(&row, Field::Boolean(true), Field::Null)] + ); + assert_eq!( + feed(&mut p, VALUES_PORT, delete(inner(Field::Binary(vec![1])))).unwrap(), + vec![transition(&row, Field::Null, Field::Boolean(false))] + ); +} + +#[test] +fn volatile_operand_is_rejected_even_inside_a_scalar_expression() { + let runtime = crate::tests::utils::create_test_runtime(); + for sql in [ + "SELECT * FROM rows WHERE NOW()", + "SELECT * FROM rows WHERE NOW() IS NULL", + ] { + let select = crate::tests::utils::get_select(sql).unwrap(); + let operand = runtime + .block_on(ExpressionBuilder::new(2, runtime.clone()).build( + false, + &select.selection.unwrap(), + &schema(), + &[], + )) + .unwrap(); + let error = MembershipProcessor::new(schema(), operand, false).unwrap_err(); + assert!(error + .to_string() + .contains("deterministic scalar expressions")); + } +} + +#[test] +fn updating_into_an_existing_outer_row_preserves_multiplicity() { + let mut p = processor(false); + let old = outer(1, Field::Int(3)); + let new = outer(2, Field::Int(7)); + feed(&mut p, OUTER_PORT, insert(old.clone())).unwrap(); + feed(&mut p, OUTER_PORT, insert(new.clone())).unwrap(); + feed( + &mut p, + OUTER_PORT, + Operation::Update { + old, + new: new.clone(), + }, + ) + .unwrap(); + assert_eq!(p.retained(&new).unwrap().copies, 2); + assert_eq!( + feed(&mut p, VALUES_PORT, insert(inner(Field::Int(7)))).unwrap(), + vec![transition(&new, Field::Boolean(false), Field::Boolean(true)); 2] + ); + feed(&mut p, OUTER_PORT, delete(new.clone())).unwrap(); + feed(&mut p, OUTER_PORT, delete(new)).unwrap(); + assert!(p.rows.is_empty()); +} + +fn representation_sensitive_operand(sql: &str, typ: FieldType) -> MembershipProcessor { + let runtime = crate::tests::utils::create_test_runtime(); + let mut schema = schema(); + schema.fields[1].typ = typ; + let select = crate::tests::utils::get_select(sql).unwrap(); + let operand = runtime + .block_on(ExpressionBuilder::new(2, runtime.clone()).build( + false, + &select.selection.unwrap(), + &schema, + &[], + )) + .unwrap(); + MembershipProcessor::new(schema, operand, false).unwrap() +} + +fn assert_distinct_representations( + mut p: MembershipProcessor, + first: Record, + second: Record, + matching: Field, +) { + assert_eq!( + first, second, + "this regression requires semantic Record equality" + ); + assert_ne!(row_key(&first).unwrap(), row_key(&second).unwrap()); + feed(&mut p, VALUES_PORT, insert(inner(matching.clone()))).unwrap(); + assert_eq!( + feed(&mut p, OUTER_PORT, insert(first.clone())).unwrap(), + vec![insert(marked(&first, Field::Boolean(true)))] + ); + assert_eq!( + feed(&mut p, OUTER_PORT, insert(second.clone())).unwrap(), + vec![insert(marked(&second, Field::Boolean(false)))] + ); + assert_eq!(p.rows.len(), 2); + let changes = feed(&mut p, VALUES_PORT, delete(inner(matching))).unwrap(); + assert_eq!( + changes, + vec![transition( + &first, + Field::Boolean(true), + Field::Boolean(false) + )] + ); + // Compare encoded records too: ordinary equality hides the very difference + // which this test is intended to protect. + match &changes[0] { + Operation::Update { new, .. } => assert_eq!( + row_key(new).unwrap(), + row_key(&marked(&first, Field::Boolean(false))).unwrap() + ), + _ => panic!("expected membership update"), + } + let deletion = feed(&mut p, OUTER_PORT, delete(second.clone())).unwrap(); + assert_eq!( + deletion, + vec![delete(marked(&second, Field::Boolean(false)))] + ); + assert_eq!(p.rows.len(), 1); + assert_eq!(p.retained(&first).unwrap().copies, 1); + feed(&mut p, OUTER_PORT, delete(first)).unwrap(); + assert!(p.rows.is_empty()); +} + +#[test] +fn equal_decimal_values_keep_distinct_scale_sensitive_operands() { + use dozer_types::rust_decimal::Decimal; + assert_distinct_representations( + representation_sensitive_operand( + "SELECT * FROM rows WHERE CAST(value AS STRING)", + FieldType::Decimal, + ), + outer(1, Field::Decimal(Decimal::new(10, 1))), + outer(1, Field::Decimal(Decimal::new(100, 2))), + Field::String("1.0".into()), + ); +} + +#[test] +fn equal_timestamp_instants_keep_distinct_local_hour_operands() { + use dozer_types::chrono::DateTime; + assert_distinct_representations( + representation_sensitive_operand( + "SELECT * FROM rows WHERE EXTRACT(HOUR FROM value)", + FieldType::Timestamp, + ), + outer( + 1, + Field::Timestamp(DateTime::parse_from_rfc3339("2024-01-01T00:00:00+00:00").unwrap()), + ), + outer( + 1, + Field::Timestamp(DateTime::parse_from_rfc3339("2024-01-01T01:00:00+01:00").unwrap()), + ), + Field::Int(0), + ); +} + +#[test] +fn factory_validates_width_and_builds_a_nullable_membership_column() { + use dozer_sql_expression::sqlparser::ast::Ident; + let runtime = crate::tests::utils::create_test_runtime(); + let factory = MembershipProcessorFactory { + id: "membership".into(), + operand: Expr::Identifier(Ident::new("value")), + negated: false, + column: "membership_result".into(), + udfs: vec![], + runtime: runtime.clone(), + }; + let mut values = schema(); + values.fields.remove(0); + values.primary_index.clear(); + let inputs = HashMap::from([(OUTER_PORT, schema()), (VALUES_PORT, values)]); + runtime.block_on(async { + let output = factory + .get_output_schema(&DEFAULT_PORT_HANDLE, &inputs) + .await + .unwrap(); + assert_eq!(&output.fields[..2], schema().fields); + assert_eq!(output.primary_index, vec![0]); + assert_eq!(output.fields[2].typ, FieldType::Boolean); + assert!(output.fields[2].nullable); + let mut processor = factory + .build(inputs.clone(), HashMap::new(), EventHub::new(1)) + .await + .unwrap(); + let mut captured = Output::default(); + processor + .process( + TableOperation { + id: None, + port: OUTER_PORT, + op: insert(outer(1, Field::Int(7))), + }, + &mut captured, + ) + .unwrap(); + assert_eq!( + captured.0[0].op, + insert(marked(&outer(1, Field::Int(7)), Field::Boolean(false))) + ); + let mut invalid = inputs; + invalid.insert(VALUES_PORT, schema()); + assert!(factory + .get_output_schema(&DEFAULT_PORT_HANDLE, &invalid) + .await + .is_err()); + invalid.remove(&VALUES_PORT); + assert!(factory + .get_output_schema(&DEFAULT_PORT_HANDLE, &invalid) + .await + .is_err()); + }); +} + +fn apply(multiset: &mut Vec, operation: Operation) { + match operation { + Operation::Insert { new } => multiset.push(new), + Operation::Delete { old } => { + let position = multiset.iter().position(|row| row == &old).unwrap(); + multiset.remove(position); + } + Operation::Update { old, new } => { + apply(multiset, delete(old)); + apply(multiset, insert(new)); + } + Operation::BatchInsert { new } => multiset.extend(new), + } +} + +fn frequencies(rows: impl IntoIterator) -> HashMap { + let mut result = HashMap::new(); + for row in rows { + *result.entry(row).or_default() += 1; + } + result +} + +#[test] +fn event_stream_matches_an_independent_snapshot_multiset() { + // Reference calculation intentionally uses only integer/NULL snapshot rows, + // with no cached match counts or the processor's comparison helpers. + for negated in [false, true] { + let mut p = processor(negated); + let mut outer_rows: Vec = Vec::new(); + let mut inner_rows: Vec = Vec::new(); + let mut observed = Vec::new(); + let mut random = 0x1665_u64; + for step in 0..600 { + random = random.wrapping_mul(6364136223846793005).wrapping_add(1); + let port = if random & 8 == 0 { + OUTER_PORT + } else { + VALUES_PORT + }; + let rows = if port == OUTER_PORT { + &mut outer_rows + } else { + &mut inner_rows + }; + let value = match (random >> 8) % 5 { + 0 => Field::Null, + value => Field::Int(value as i64), + }; + let new = if port == OUTER_PORT { + outer(((random >> 16) % 4) as i64, value) + } else { + inner(value) + }; + let operation = match (random >> 24) % 4 { + 0 if !rows.is_empty() => delete(rows[(random as usize >> 32) % rows.len()].clone()), + 1 if !rows.is_empty() => Operation::Update { + old: rows[(random as usize >> 32) % rows.len()].clone(), + new, + }, + 2 => Operation::BatchInsert { + new: vec![new.clone(), new], + }, + _ => insert(new), + }; + apply(rows, operation.clone()); + for operation in feed(&mut p, port, operation).unwrap() { + apply(&mut observed, operation); + } + let expected = outer_rows.iter().map(|row| { + let value = &row.values[1]; + let found = value != &Field::Null + && inner_rows.iter().any(|inner| &inner.values[0] == value); + let unknown = !inner_rows.is_empty() + && !found + && (value == &Field::Null + || inner_rows + .iter() + .any(|inner| inner.values[0] == Field::Null)); + let mark = if unknown { + Field::Null + } else { + Field::Boolean(found ^ negated) + }; + marked(row, mark) + }); + assert_eq!( + frequencies(observed.clone()), + frequencies(expected), + "step={step}, negated={negated}" + ); + } + } +} diff --git a/dozer-sql/src/selection/mod.rs b/dozer-sql/src/selection/mod.rs index 9a12ba12cc..699dd66b5f 100644 --- a/dozer-sql/src/selection/mod.rs +++ b/dozer-sql/src/selection/mod.rs @@ -1,2 +1,4 @@ pub mod factory; +pub(crate) mod membership; pub mod processor; +pub(crate) mod replay; diff --git a/dozer-sql/src/selection/replay.rs b/dozer-sql/src/selection/replay.rs new file mode 100644 index 0000000000..16cab361df --- /dev/null +++ b/dozer-sql/src/selection/replay.rs @@ -0,0 +1,41 @@ +use std::sync::{Arc, Mutex}; + +use dozer_sql_expression::execution::Expression; + +use crate::errors::PipelineError; + +/// A downstream stateful consumer can require an already planned relation to +/// reproduce the same values when an old row is deleted or updated. +#[derive(Debug, Clone, Default)] +pub(crate) struct ReplayRequirement(Arc>); + +#[derive(Debug, Default)] +struct ReplayPolicy { + required: bool, + unsupported: Option, +} + +impl ReplayRequirement { + pub(crate) fn require(&self) { + self.0.lock().unwrap().required = true; + } + + pub(crate) fn is_required(&self) -> bool { + self.0.lock().unwrap().required + } + + pub(crate) fn reject_when_required(&self, reason: String) { + self.0.lock().unwrap().unsupported.get_or_insert(reason); + } + + pub(crate) fn validate(&self, expression: &Expression) -> Result<(), PipelineError> { + let policy = self.0.lock().unwrap(); + if policy.required { + if let Some(reason) = &policy.unsupported { + return Err(PipelineError::InvalidQuery(reason.clone())); + } + super::membership::validate_operand(expression)?; + } + Ok(()) + } +} diff --git a/dozer-sql/src/tests/in_subquery.rs b/dozer-sql/src/tests/in_subquery.rs new file mode 100644 index 0000000000..e530ae5092 --- /dev/null +++ b/dozer-sql/src/tests/in_subquery.rs @@ -0,0 +1,1375 @@ +use std::{ + collections::{HashMap, VecDeque}, + sync::Arc, +}; + +use dozer_core::{ + app::{App, AppPipeline, PipelineFlags}, + appsource::{AppSourceManager, AppSourceMappings}, + channels::ProcessorChannelForwarder, + dag_schemas::{DagHaveSchemas, DagSchemas}, + daggy::{ + petgraph::visit::{EdgeRef, Topo}, + NodeIndex, Walker, + }, + event::EventHub, + node::{OutputPortDef, OutputPortType, PortHandle, Processor, Source, SourceFactory}, + NodeKind, DEFAULT_PORT_HANDLE, +}; +use dozer_types::{ + errors::internal::BoxedError, + types::{ + Field, FieldDefinition, FieldType, Operation, Record, Schema, SourceDefinition, + TableOperation, + }, +}; +use tokio::runtime::Runtime; + +use crate::builder::statement_to_pipeline; + +use super::{builder_test::TestSinkFactory, utils::create_test_runtime}; + +#[derive(Debug)] +struct TableSources(Vec<(String, Schema)>); + +impl SourceFactory for TableSources { + fn get_output_schema(&self, port: &PortHandle) -> Result { + Ok(self.0[usize::from(*port)].1.clone()) + } + + fn get_output_port_name(&self, port: &PortHandle) -> String { + self.0[usize::from(*port)].0.clone() + } + + fn get_output_ports(&self) -> Vec { + (0..self.0.len()) + .map(|port| OutputPortDef::new(port as PortHandle, OutputPortType::Stateless)) + .collect() + } + + fn build( + &self, + _schemas: HashMap, + _events: EventHub, + _state: Option>, + ) -> Result, BoxedError> { + unreachable!("the test supplies source operations directly") + } +} + +#[derive(Default)] +struct Output(Vec); + +impl ProcessorChannelForwarder for Output { + fn send(&mut self, operation: TableOperation) { + self.0.push(operation); + } +} + +type Endpoint = (NodeIndex, PortHandle); + +struct Pipeline { + processors: HashMap>, + routes: HashMap>, + sources: HashMap, + sink: NodeIndex, + schema: Schema, + rows: Vec, + _runtime: Arc, +} + +impl Pipeline { + fn new(sql: &str, tables: Vec<(&str, Schema)>) -> Self { + Self::try_new(sql, tables).unwrap() + } + + fn try_new(sql: &str, tables: Vec<(&str, Schema)>) -> Result { + Self::try_new_with_flags(sql, tables, PipelineFlags::default()) + } + + fn try_new_with_flags( + sql: &str, + tables: Vec<(&str, Schema)>, + flags: PipelineFlags, + ) -> Result { + let runtime = create_test_runtime(); + let mut pipeline = AppPipeline::new(flags); + let context = statement_to_pipeline( + sql, + &mut pipeline, + Some("result".into()), + vec![], + runtime.clone(), + )?; + let result = &context.output_tables_map["result"]; + pipeline.add_sink( + Box::new(TestSinkFactory::new(vec![DEFAULT_PORT_HANDLE])), + "capture".into(), + ); + pipeline.connect_nodes( + result.node.clone(), + result.port, + "capture".into(), + DEFAULT_PORT_HANDLE, + ); + + let mappings: HashMap<_, _> = tables + .iter() + .enumerate() + .map(|(port, (name, _))| ((*name).to_owned(), port as PortHandle)) + .collect(); + let mut sources = AppSourceManager::new(); + sources + .add( + Box::new(TableSources( + tables + .into_iter() + .map(|(name, schema)| (name.to_owned(), schema)) + .collect(), + )), + AppSourceMappings::new("fixtures".into(), mappings.clone()), + ) + .unwrap(); + let mut app = App::new(sources); + app.add_pipeline(pipeline); + let dag = runtime.block_on(DagSchemas::new(app.into_dag()?))?; + + let mut processors = HashMap::new(); + let mut routes: HashMap> = HashMap::new(); + let mut source_node = None; + let mut sink = None; + let mut schema = None; + for index in Topo::new(dag.graph()).iter(dag.graph()) { + match &dag.graph()[index].kind { + NodeKind::Source(_) => source_node = Some(index), + NodeKind::Processor(factory) => { + let processor = runtime.block_on(factory.build( + dag.get_node_input_schemas(index), + dag.get_node_output_schemas(index), + EventHub::new(1), + ))?; + processors.insert(index, processor); + } + NodeKind::Sink(_) => { + sink = Some(index); + schema = Some(dag.get_node_input_schemas(index)[&DEFAULT_PORT_HANDLE].clone()); + } + } + } + for edge in dag.graph().graph().edge_references() { + routes + .entry((edge.source(), edge.weight().output_port)) + .or_default() + .push((edge.target(), edge.weight().input_port)); + } + let source_node = source_node.unwrap(); + Ok(Self { + processors, + routes, + sources: mappings + .into_iter() + .map(|(name, port)| (name, (source_node, port))) + .collect(), + sink: sink.unwrap(), + schema: schema.unwrap(), + rows: vec![], + _runtime: runtime, + }) + } + + fn two_tables(sql: &str) -> Self { + Self::new( + sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ) + } + + fn feed(&mut self, table: &str, operation: Operation) -> Vec { + // Every operator is built by the real SQL planner. Only the transport + // is synchronous, so each assertion observes a drained event queue. + let (source, port) = self.sources[table]; + let mut queue = VecDeque::from([(source, TableOperation::without_id(operation, port))]); + let mut emitted = Vec::new(); + while let Some((from, operation)) = queue.pop_front() { + for &(to, port) in self.routes.get(&(from, operation.port)).unwrap() { + let operation = TableOperation { + id: operation.id, + port, + op: operation.op.clone(), + }; + if to == self.sink { + apply(&mut self.rows, &operation.op); + emitted.push(operation.op); + } else { + let mut output = Output::default(); + self.processors + .get_mut(&to) + .unwrap() + .process(operation, &mut output) + .unwrap(); + queue.extend(output.0.into_iter().map(|operation| (to, operation))); + } + } + } + emitted + } + + fn assert_rows(&self, expected: Vec) { + fn counts(rows: &[Record]) -> HashMap { + let mut counts = HashMap::new(); + for row in rows { + *counts.entry(row.clone()).or_default() += 1; + } + counts + } + assert_eq!(counts(&self.rows), counts(&expected)); + assert!(self + .rows + .iter() + .all(|row| row.values.len() == self.schema.fields.len())); + } +} + +fn table_schema(name: &str, with_id: bool) -> Schema { + let source = SourceDefinition::Table { + connection: "fixtures".into(), + name: name.into(), + }; + let mut fields = vec![]; + if with_id { + fields.push(FieldDefinition::new( + "id".into(), + FieldType::Int, + false, + source.clone(), + )); + } + fields.push(FieldDefinition::new( + "value".into(), + FieldType::Int, + true, + source, + )); + Schema { + fields, + primary_index: vec![], + } +} + +fn outer(id: i64, value: Field) -> Record { + Record::new(vec![Field::Int(id), value]) +} + +fn inner(value: Field) -> Record { + Record::new(vec![value]) +} + +fn insert(new: Record) -> Operation { + Operation::Insert { new } +} + +fn delete(old: Record) -> Operation { + Operation::Delete { old } +} + +fn apply(rows: &mut Vec, operation: &Operation) { + match operation { + Operation::Insert { new } => rows.push(new.clone()), + Operation::Delete { old } => { + let index = rows + .iter() + .position(|row| row == old) + .expect("the pipeline deleted a row that was never emitted"); + rows.remove(index); + } + Operation::Update { old, new } => { + apply(rows, &delete(old.clone())); + apply(rows, &insert(new.clone())); + } + Operation::BatchInsert { new } => rows.extend(new.iter().cloned()), + } +} + +#[test] +fn late_inner_changes_and_outer_changes_reach_the_query_output() { + let mut pipeline = Pipeline::two_tables( + "SELECT * FROM outer_rows WHERE value IN (SELECT value FROM inner_rows)", + ); + assert_eq!( + pipeline + .schema + .fields + .iter() + .map(|field| field.name.as_str()) + .collect::>(), + ["id", "value"] + ); + let first = outer(1, Field::Int(7)); + let second = outer(2, Field::Int(9)); + pipeline.feed("outer_rows", insert(first.clone())); + pipeline.feed("outer_rows", insert(second.clone())); + pipeline.assert_rows(vec![]); + assert_eq!( + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))), + vec![insert(first.clone())] + ); + let changed = outer(1, Field::Int(9)); + assert_eq!( + pipeline.feed( + "outer_rows", + Operation::Update { + old: first.clone(), + new: changed.clone() + } + ), + vec![delete(first)] + ); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![changed.clone(), second.clone()]); + pipeline.feed("outer_rows", delete(second)); + pipeline.assert_rows(vec![changed]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn duplicate_and_null_membership_transitions_preserve_sql_rows() { + for negated in [false, true] { + let mut pipeline = Pipeline::two_tables(&format!( + "SELECT * FROM outer_rows WHERE value {}IN (SELECT value FROM inner_rows)", + if negated { "NOT " } else { "" } + )); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + let null = outer(3, Field::Null); + let rows = vec![seven.clone(), seven.clone(), nine.clone(), null]; + pipeline.feed("outer_rows", Operation::BatchInsert { new: rows.clone() }); + pipeline.assert_rows(if negated { rows.clone() } else { vec![] }); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![ + inner(Field::Int(7)), + inner(Field::Int(7)), + inner(Field::Null), + ], + }, + ); + pipeline.assert_rows(if negated { + vec![] + } else { + vec![seven.clone(), seven.clone()] + }); + assert!(pipeline + .feed("inner_rows", delete(inner(Field::Int(7)))) + .is_empty()); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", delete(inner(Field::Null))); + pipeline.assert_rows(if negated { rows } else { vec![] }); + pipeline.feed("inner_rows", insert(inner(Field::Int(9)))); + pipeline.assert_rows(if negated { + vec![seven.clone(), seven.clone()] + } else { + vec![nine.clone()] + }); + pipeline.feed("outer_rows", delete(seven.clone())); + pipeline.assert_rows(if negated { vec![seven] } else { vec![nine] }); + } +} + +#[test] +fn projection_receives_only_visible_columns() { + let mut pipeline = Pipeline::two_tables( + "SELECT id FROM outer_rows WHERE value IN (SELECT value FROM inner_rows)", + ); + assert_eq!(pipeline.schema.fields.len(), 1); + assert_eq!(pipeline.schema.fields[0].name, "id"); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.feed("outer_rows", insert(outer(4, Field::Int(7)))); + pipeline.assert_rows(vec![Record::new(vec![Field::Int(4)])]); +} + +#[test] +fn qualified_operands_and_filtered_inner_rows_use_their_own_schemas() { + let mut pipeline = Pipeline::two_tables( + "SELECT o.id FROM outer_rows o WHERE o.value IN \ + (SELECT i.value FROM inner_rows i WHERE i.value > 5)", + ); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![outer(1, Field::Int(7)), outer(2, Field::Int(4))], + }, + ); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![inner(Field::Int(7)), inner(Field::Int(4))], + }, + ); + pipeline.assert_rows(vec![Record::new(vec![Field::Int(1)])]); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(4)), + }, + ); + pipeline.assert_rows(vec![]); +} + +#[test] +fn outer_cte_membership_uses_the_visible_relation_name_or_alias() { + for subquery in ["SELECT c.value FROM c", "SELECT x.value FROM c AS x"] { + let mut pipeline = Pipeline::two_tables(&format!( + "WITH c AS (SELECT value FROM inner_rows) \ + SELECT * FROM outer_rows WHERE value IN ({subquery})" + )); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![seven.clone(), nine.clone(), outer(3, Field::Null)], + }, + ); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven]); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); + assert_eq!( + pipeline + .schema + .fields + .iter() + .map(|field| field.name.as_str()) + .collect::>(), + ["id", "value"] + ); + } +} + +#[test] +fn cte_membership_does_not_expose_original_or_hidden_qualifiers() { + for subquery in [ + "SELECT inner_rows.value FROM c", + "SELECT c.value FROM c AS x", + "SELECT inner_rows.value FROM c AS x", + ] { + let sql = format!( + "WITH c AS (SELECT value FROM inner_rows) \ + SELECT * FROM outer_rows WHERE value IN ({subquery})" + ); + let result = Pipeline::try_new( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ); + assert!( + result.is_err(), + "CTE accepted an out-of-scope qualifier: {subquery}" + ); + } +} + +#[test] +fn volatile_inner_filters_are_rejected_including_cte_dependencies() { + use dozer_sql_expression::sqlparser::{dialect::DozerDialect, parser::Parser}; + + for sql in [ + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value FROM inner_rows WHERE NOW() IS NOT NULL)", + "WITH c AS (SELECT value FROM inner_rows WHERE NOW() IS NOT NULL) \ + SELECT * FROM outer_rows WHERE value IN (SELECT c.value FROM c)", + "WITH c AS (SELECT value FROM inner_rows WHERE NOW() IS NOT NULL), \ + d AS (SELECT c.value FROM c) \ + SELECT * FROM outer_rows WHERE value IN (SELECT d.value FROM d)", + "WITH c AS (SELECT NOW() AS value FROM inner_rows) \ + SELECT * FROM outer_rows WHERE value IN (SELECT c.value FROM c)", + "SELECT * FROM outer_rows WHERE \ + value IN (SELECT value FROM inner_rows) AND NOW() IS NOT NULL", + ] { + Parser::parse_sql(&DozerDialect {}, sql) + .unwrap_or_else(|error| panic!("test must reach planner validation: {sql}: {error}")); + let result = Pipeline::try_new( + sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ); + assert!( + result.is_err(), + "volatile membership input was accepted: {sql}" + ); + } +} + +#[test] +fn ordinary_queries_keep_volatile_filters_without_membership() { + let mut pipeline = Pipeline::new( + "SELECT * FROM outer_rows WHERE NOW() IS NOT NULL", + vec![("outer_rows", table_schema("outer_rows", true))], + ); + let row = outer(1, Field::Int(7)); + pipeline.feed("outer_rows", insert(row.clone())); + pipeline.assert_rows(vec![row.clone()]); + pipeline.feed("outer_rows", delete(row)); + pipeline.assert_rows(vec![]); +} + +#[test] +fn membership_inside_a_compound_predicate_keeps_the_other_branch() { + let mut pipeline = Pipeline::two_tables( + "SELECT * FROM outer_rows WHERE id = 99 OR \ + (value IN (SELECT value FROM inner_rows) AND value > 5)", + ); + let always = outer(99, Field::Int(1)); + let member = outer(1, Field::Int(7)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![always.clone(), member.clone(), outer(2, Field::Int(3))], + }, + ); + pipeline.assert_rows(vec![always.clone()]); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![inner(Field::Int(7)), inner(Field::Int(3))], + }, + ); + pipeline.assert_rows(vec![always.clone(), member]); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![always]); +} + +#[test] +fn negated_compounds_preserve_unknown_membership() { + for predicate in [ + "NOT (value IN (SELECT value FROM inner_rows) OR 1 = 2)", + "NOT (value IN (SELECT value FROM inner_rows) AND 1 = 1)", + ] { + let mut pipeline = + Pipeline::two_tables(&format!("SELECT * FROM outer_rows WHERE {predicate}")); + let row = outer(1, Field::Int(7)); + pipeline.feed("outer_rows", insert(row.clone())); + pipeline.assert_rows(vec![row.clone()]); + pipeline.feed("inner_rows", insert(inner(Field::Null))); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", delete(inner(Field::Null))); + pipeline.assert_rows(vec![row]); + } +} + +#[test] +fn unsupported_subqueries_fail_during_pipeline_construction() { + for subquery in [ + "SELECT value, value FROM inner_rows", + "SELECT DISTINCT value FROM inner_rows", + ] { + let sql = format!("SELECT * FROM outer_rows WHERE value IN ({subquery})"); + let result = Pipeline::try_new( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ); + assert!( + result.is_err(), + "unsupported subquery was accepted: {subquery}" + ); + } +} + +#[test] +fn inner_with_shadows_an_outer_cte_without_changing_sibling_membership() { + let mut pipeline = Pipeline::new( + "WITH candidates AS (SELECT value FROM other_rows) \ + SELECT * FROM outer_rows WHERE \ + value IN (WITH candidates AS (SELECT value FROM inner_rows) \ + SELECT candidates.value FROM candidates) \ + AND value IN (SELECT candidates.value FROM candidates)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![ + seven.clone(), + seven.clone(), + nine.clone(), + outer(3, Field::Null), + ], + }, + ); + pipeline.feed("other_rows", insert(inner(Field::Int(9)))); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", insert(inner(Field::Int(9)))); + pipeline.assert_rows(vec![nine]); + pipeline.feed( + "other_rows", + Operation::Update { + old: inner(Field::Int(9)), + new: inner(Field::Int(7)), + }, + ); + pipeline.assert_rows(vec![seven.clone(), seven]); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![]); + assert_eq!(pipeline.schema.fields.len(), 2); +} + +#[test] +fn a_sibling_membership_cannot_read_a_previous_subquery_private_cte() { + let result = Pipeline::try_new( + "SELECT * FROM outer_rows WHERE \ + value IN (WITH private_values AS (SELECT value FROM inner_rows) \ + SELECT value FROM private_values) \ + AND value IN (SELECT value FROM private_values)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ); + assert!( + result.is_err(), + "a sibling query exposed the inner WITH binding" + ); +} + +#[test] +fn nested_and_multiple_membership_queries_track_nulls_and_duplicates() { + let mut pipeline = Pipeline::new( + "SELECT * FROM outer_rows WHERE \ + value IN (SELECT value FROM inner_rows WHERE value IN \ + (SELECT value FROM allowed_rows)) \ + AND value NOT IN (SELECT value FROM blocked_rows)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("allowed_rows", table_schema("allowed_rows", false)), + ("blocked_rows", table_schema("blocked_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![ + seven.clone(), + seven.clone(), + nine.clone(), + outer(3, Field::Null), + ], + }, + ); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![ + inner(Field::Int(7)), + inner(Field::Int(7)), + inner(Field::Int(9)), + inner(Field::Null), + ], + }, + ); + pipeline.assert_rows(vec![]); + pipeline.feed("allowed_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven.clone()]); + pipeline.feed("blocked_rows", insert(inner(Field::Null))); + pipeline.assert_rows(vec![]); + pipeline.feed("blocked_rows", insert(inner(Field::Int(7)))); + pipeline.feed("blocked_rows", delete(inner(Field::Null))); + pipeline.assert_rows(vec![]); + pipeline.feed("blocked_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven]); + pipeline.feed( + "allowed_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.feed("allowed_rows", insert(inner(Field::Null))); + pipeline.assert_rows(vec![]); + assert_eq!(pipeline.schema.fields.len(), 2); +} + +#[test] +fn computed_inner_projection_changes_membership_after_updates() { + let mut pipeline = Pipeline::two_tables( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value + 1 AS next_value FROM inner_rows)", + ); + let eight = outer(1, Field::Int(8)); + let ten = outer(2, Field::Int(10)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![eight.clone(), ten.clone(), outer(3, Field::Null)], + }, + ); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![eight]); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![ten.clone()]); + pipeline.feed("inner_rows", insert(inner(Field::Null))); + pipeline.assert_rows(vec![ten]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn union_membership_retains_matches_until_both_branches_remove_them() { + for quantifier in ["ALL", "DISTINCT"] { + let mut pipeline = Pipeline::new( + &format!( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value FROM inner_rows UNION {quantifier} \ + SELECT value FROM other_rows)" + ), + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![ + seven.clone(), + seven.clone(), + nine.clone(), + outer(3, Field::Null), + ], + }, + ); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![ + inner(Field::Int(7)), + inner(Field::Int(7)), + inner(Field::Null), + ], + }, + ); + pipeline.feed( + "other_rows", + Operation::BatchInsert { + new: vec![inner(Field::Int(7)), inner(Field::Int(9))], + }, + ); + pipeline.assert_rows(vec![seven.clone(), seven.clone(), nine.clone()]); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven, nine.clone()]); + pipeline.feed( + "other_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine.clone()]); + pipeline.feed("other_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![nine]); + pipeline.feed("other_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); + } +} + +#[test] +fn grouped_having_membership_changes_when_a_group_crosses_its_threshold() { + let mut pipeline = Pipeline::two_tables( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value FROM inner_rows GROUP BY value HAVING COUNT(*) > 1)", + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![seven.clone(), seven.clone(), nine.clone()], + }, + ); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven.clone()]); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![inner(Field::Int(9)), inner(Field::Int(9))], + }, + ); + pipeline.assert_rows(vec![seven.clone(), seven, nine.clone()]); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine.clone()]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![nine]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn joined_inner_rows_keep_membership_until_the_last_match_disappears() { + let mut pipeline = Pipeline::new( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT i.value FROM inner_rows i JOIN other_rows j ON i.value = j.value)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![seven.clone(), seven.clone(), nine.clone()], + }, + ); + pipeline.feed( + "inner_rows", + Operation::BatchInsert { + new: vec![ + inner(Field::Int(7)), + inner(Field::Int(7)), + inner(Field::Int(9)), + ], + }, + ); + pipeline.assert_rows(vec![]); + pipeline.feed("other_rows", insert(inner(Field::Int(7)))); + pipeline.feed("other_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven.clone()]); + pipeline.feed("other_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone(), seven]); + pipeline.feed( + "other_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn expanded_inner_query_dependencies_reject_volatile_evaluation() { + use dozer_sql_expression::sqlparser::{dialect::DozerDialect, parser::Parser}; + + for subquery in [ + "WITH candidate AS (SELECT value FROM inner_rows WHERE NOW() IS NOT NULL) \ + SELECT value FROM candidate", + "SELECT value FROM inner_rows WHERE NOW() IS NOT NULL \ + UNION ALL SELECT value FROM other_rows", + "SELECT value FROM inner_rows \ + UNION DISTINCT SELECT value FROM other_rows WHERE NOW() IS NOT NULL", + "WITH candidate AS (SELECT value FROM inner_rows WHERE NOW() IS NOT NULL) \ + SELECT c.value FROM candidate c JOIN other_rows j ON c.value = j.value", + "WITH candidate AS (SELECT value FROM other_rows WHERE NOW() IS NOT NULL) \ + SELECT i.value FROM inner_rows i JOIN candidate c ON i.value = c.value", + "SELECT i.value FROM inner_rows i JOIN other_rows j ON i.value = j.value \ + WHERE NOW() IS NOT NULL", + ] { + let sql = format!("SELECT * FROM outer_rows WHERE value IN ({subquery})"); + Parser::parse_sql(&DozerDialect {}, &sql).unwrap(); + let result = Pipeline::try_new( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let error = match result { + Ok(_) => panic!("volatile membership dependency was accepted: {subquery}"), + Err(error) => error, + }; + assert!( + error.to_string().contains("deterministic"), + "expected the replay guard to reject {subquery}, got {error}" + ); + } +} + +#[test] +fn ancestor_cte_and_derived_modifiers_cannot_bypass_membership_validation() { + use dozer_sql_expression::sqlparser::{dialect::DozerDialect, parser::Parser}; + + for (subquery, expected) in [ + ("SELECT TOP (1) value FROM inner_rows", "SELECT modifiers"), + ( + "SELECT value FROM inner_rows FETCH FIRST 1 ROW ONLY", + "FETCH", + ), + ] { + for sql in [ + format!("SELECT * FROM outer_rows WHERE value IN ({subquery})"), + format!( + "WITH candidate AS ({subquery}) \ + SELECT * FROM outer_rows WHERE value IN (SELECT value FROM candidate)" + ), + format!( + "WITH candidate AS \ + (SELECT value FROM inner_rows UNION ALL ({subquery})) \ + SELECT * FROM outer_rows WHERE value IN (SELECT value FROM candidate)" + ), + format!( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value FROM ({subquery}) derived_rows)" + ), + ] { + Parser::parse_sql(&DozerDialect {}, &sql).unwrap(); + let error = match Pipeline::try_new( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ) { + Ok(_) => panic!("unsupported dependency modifier was accepted: {sql}"), + Err(error) => error, + }; + assert!(error.to_string().contains(expected), "{sql}: {error}"); + } + } +} + +#[test] +fn unsupported_aggregate_shapes_are_rejected_before_processing() { + use dozer_sql_expression::sqlparser::{dialect::DozerDialect, parser::Parser}; + + for (sql, expected) in [ + ( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT COUNT(*) FROM inner_rows)", + "without GROUP BY", + ), + ( + "WITH candidate AS (SELECT COUNT(*) AS value FROM inner_rows) \ + SELECT * FROM outer_rows WHERE value IN (SELECT value FROM candidate)", + "without GROUP BY", + ), + ( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT COUNT(DISTINCT value) FROM inner_rows GROUP BY value)", + "function modifiers", + ), + ( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT COUNT(*) OVER () FROM inner_rows GROUP BY value)", + "function modifiers", + ), + ] { + Parser::parse_sql(&DozerDialect {}, sql).unwrap(); + let error = match Pipeline::try_new( + sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ) { + Ok(_) => panic!("unsupported aggregate was accepted: {sql}"), + Err(error) => error, + }; + assert!(error.to_string().contains(expected), "{sql}: {error}"); + } +} + +#[test] +fn parenthesized_select_and_union_branches_keep_inner_with_bindings() { + for (subquery, union) in [ + ( + "WITH candidate AS (SELECT value FROM inner_rows) \ + (SELECT candidate.value FROM candidate)", + false, + ), + ( + "SELECT value FROM inner_rows UNION ALL \ + (WITH candidate AS (SELECT value FROM other_rows) \ + SELECT candidate.value FROM candidate)", + true, + ), + ] { + let mut pipeline = Pipeline::new( + &format!("SELECT * FROM outer_rows WHERE value IN ({subquery})"), + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![seven.clone(), nine.clone()], + }, + ); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven.clone()]); + if union { + pipeline.feed("other_rows", insert(inner(Field::Int(9)))); + pipeline.assert_rows(vec![seven, nine.clone()]); + } + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![nine.clone()]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + if union { + pipeline.assert_rows(vec![nine]); + pipeline.feed("other_rows", delete(inner(Field::Int(9)))); + } + pipeline.assert_rows(vec![]); + } +} + +#[test] +fn cte_join_without_alias_uses_its_visible_cte_name() { + let mut pipeline = Pipeline::new( + "WITH candidate AS (SELECT value FROM inner_rows) \ + SELECT * FROM outer_rows WHERE value IN \ + (SELECT candidate.value FROM candidate JOIN other_rows \ + ON candidate.value = other_rows.value)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + ); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![seven.clone(), nine.clone()], + }, + ); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![]); + pipeline.feed("other_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![seven]); + pipeline.feed( + "other_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![]); + pipeline.feed("inner_rows", insert(inner(Field::Int(9)))); + pipeline.assert_rows(vec![nine]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn grouped_count_projection_keeps_shared_counts_until_the_last_group_changes() { + let mut pipeline = Pipeline::two_tables( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT COUNT(*) FROM inner_rows GROUP BY value)", + ); + let one = outer(1, Field::Int(1)); + let two = outer(2, Field::Int(2)); + pipeline.feed( + "outer_rows", + Operation::BatchInsert { + new: vec![one.clone(), two.clone()], + }, + ); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![one.clone()]); + pipeline.feed("inner_rows", insert(inner(Field::Int(9)))); + pipeline.assert_rows(vec![one.clone()]); + pipeline.feed("inner_rows", insert(inner(Field::Int(7)))); + pipeline.assert_rows(vec![one.clone(), two.clone()]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![two]); + pipeline.feed( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + ); + pipeline.assert_rows(vec![one.clone()]); + pipeline.feed("inner_rows", delete(inner(Field::Int(7)))); + pipeline.assert_rows(vec![one]); + pipeline.feed("inner_rows", delete(inner(Field::Int(9)))); + pipeline.assert_rows(vec![]); +} + +#[test] +fn probabilistic_union_distinct_is_rejected_for_direct_and_cte_membership() { + let mut flags = PipelineFlags::default(); + flags.enable_probabilistic_optimizations.in_sets = Some(true); + for quantifier in ["", " DISTINCT"] { + let union = + format!("SELECT value FROM inner_rows UNION{quantifier} SELECT value FROM other_rows"); + for sql in [ + format!("SELECT * FROM outer_rows WHERE value IN ({union})"), + format!( + "WITH candidate AS ({union}) SELECT * FROM outer_rows \ + WHERE value IN (SELECT value FROM candidate)" + ), + ] { + let error = match Pipeline::try_new_with_flags( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + flags.clone(), + ) { + Ok(_) => panic!("membership accepted probabilistic UNION counting: {sql}"), + Err(error) => error, + }; + assert!( + error.to_string().contains("exact UNION counting"), + "{sql}: {error}" + ); + } + } +} + +#[test] +fn probabilistic_set_flag_keeps_union_all_and_ordinary_union_distinct_working() { + let mut flags = PipelineFlags::default(); + flags.enable_probabilistic_optimizations.in_sets = Some(true); + let mut membership = Pipeline::try_new_with_flags( + "SELECT * FROM outer_rows WHERE value IN \ + (SELECT value FROM inner_rows UNION ALL SELECT value FROM other_rows)", + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + flags.clone(), + ) + .unwrap(); + let row = outer(1, Field::Int(7)); + membership.feed("outer_rows", insert(row.clone())); + membership.feed("inner_rows", insert(inner(Field::Int(7)))); + membership.assert_rows(vec![row.clone()]); + membership.feed("other_rows", insert(inner(Field::Int(7)))); + membership.feed("inner_rows", delete(inner(Field::Int(7)))); + membership.assert_rows(vec![row]); + membership.feed("other_rows", delete(inner(Field::Int(7)))); + membership.assert_rows(vec![]); + + let mut ordinary = Pipeline::try_new_with_flags( + "SELECT value FROM inner_rows UNION DISTINCT SELECT value FROM other_rows", + vec![ + ("inner_rows", table_schema("inner_rows", false)), + ("other_rows", table_schema("other_rows", false)), + ], + flags, + ) + .unwrap(); + let row = inner(Field::Int(7)); + ordinary.feed("inner_rows", insert(row.clone())); + ordinary.assert_rows(vec![row.clone()]); + ordinary.feed("other_rows", insert(row.clone())); + ordinary.assert_rows(vec![row.clone()]); + ordinary.feed("inner_rows", delete(row.clone())); + ordinary.assert_rows(vec![row.clone()]); + ordinary.feed("other_rows", delete(row)); + ordinary.assert_rows(vec![]); +} + +#[test] +fn membership_demo_with_synthetic_rows() { + let sql = "SELECT * FROM outer_rows WHERE value IN (SELECT value FROM inner_rows)"; + let mut pipeline = Pipeline::two_tables(sql); + println!("Synthetic data; real SQL planner and processors; synchronous event queue."); + println!("SQL: {sql}"); + let seven = outer(1, Field::Int(7)); + let nine = outer(2, Field::Int(9)); + let steps = [ + ("outer_rows", insert(seven.clone()), vec![]), + ("outer_rows", insert(nine.clone()), vec![]), + ("inner_rows", insert(inner(Field::Int(7))), vec![seven]), + ( + "inner_rows", + Operation::Update { + old: inner(Field::Int(7)), + new: inner(Field::Int(9)), + }, + vec![nine.clone()], + ), + ("outer_rows", delete(nine), vec![]), + ("inner_rows", delete(inner(Field::Int(9))), vec![]), + ]; + for (step, (table, operation, expected)) in steps.into_iter().enumerate() { + println!("Step {} input {table}: {operation:?}", step + 1); + let output = pipeline.feed(table, operation); + println!(" Emitted operations: {output:?}"); + println!(" Result rows: {:?}", pipeline.rows); + pipeline.assert_rows(expected); + } + println!("PASS: all six input changes produced the expected result rows."); +} + +#[test] +fn parsed_modifiers_are_rejected_instead_of_ignored() { + use crate::errors::PipelineError; + use dozer_sql_expression::sqlparser::{dialect::DozerDialect, parser::Parser}; + + for (subquery, expected) in [ + ( + "SELECT value FROM inner_rows FETCH FIRST 1 ROW ONLY", + "FETCH", + ), + ("SELECT TOP (1) value FROM inner_rows", "SELECT modifiers"), + ( + "SELECT value FROM inner_rows QUALIFY value > 5", + "SELECT modifiers", + ), + ("SELECT value FROM inner_rows(1)", "table arguments"), + ("SELECT value FROM inner_rows WITH (NOLOCK)", "hints"), + ( + "SELECT renamed FROM inner_rows AS i(renamed)", + "column aliases", + ), + ] { + let sql = format!("SELECT * FROM outer_rows WHERE value IN ({subquery})"); + Parser::parse_sql(&DozerDialect {}, &sql) + .unwrap_or_else(|error| panic!("test must reach planner validation: {sql}: {error}")); + let runtime = create_test_runtime(); + let error = statement_to_pipeline( + &sql, + &mut AppPipeline::new_with_default_flags(), + Some("result".into()), + vec![], + runtime, + ) + .unwrap_err(); + assert!( + matches!(error, PipelineError::InvalidQuery(_)), + "expected explicit planner rejection for {subquery}, got {error}" + ); + assert!(error.to_string().contains(expected), "{subquery}: {error}"); + } +} + +#[test] +fn correlated_names_are_not_rebound_to_the_inner_relation() { + for subquery in [ + "SELECT i.value FROM inner_rows i WHERE i.value = o.value", + "SELECT o.value FROM inner_rows i", + ] { + let sql = format!("SELECT * FROM outer_rows o WHERE o.value IN ({subquery})"); + let result = Pipeline::try_new( + &sql, + vec![ + ("outer_rows", table_schema("outer_rows", true)), + ("inner_rows", table_schema("inner_rows", false)), + ], + ); + assert!( + result.is_err(), + "outer reference was accepted as an inner column: {subquery}" + ); + } +} + +#[test] +fn shared_source_settles_with_either_fanout_order() { + for reverse in [false, true] { + let mut pipeline = Pipeline::new( + "SELECT * FROM rows WHERE value IN (SELECT value FROM rows)", + vec![("rows", table_schema("rows", true))], + ); + if reverse { + for targets in pipeline.routes.values_mut() { + targets.reverse(); + } + } + let first = outer(1, Field::Int(7)); + let second = outer(2, Field::Int(9)); + pipeline.feed( + "rows", + Operation::BatchInsert { + new: vec![first.clone(), second.clone(), outer(3, Field::Null)], + }, + ); + pipeline.assert_rows(vec![first.clone(), second.clone()]); + let changed = outer(1, Field::Int(10)); + pipeline.feed( + "rows", + Operation::Update { + old: first, + new: changed.clone(), + }, + ); + pipeline.assert_rows(vec![changed.clone(), second.clone()]); + pipeline.feed("rows", insert(changed.clone())); + pipeline.assert_rows(vec![changed.clone(), changed.clone(), second.clone()]); + pipeline.feed("rows", delete(second)); + pipeline.feed("rows", delete(changed.clone())); + pipeline.assert_rows(vec![changed]); + assert_eq!(pipeline.schema.fields.len(), 2); + } +} diff --git a/dozer-sql/src/tests/mod.rs b/dozer-sql/src/tests/mod.rs index 0ed1750729..1585649c13 100644 --- a/dozer-sql/src/tests/mod.rs +++ b/dozer-sql/src/tests/mod.rs @@ -1,2 +1,3 @@ mod builder_test; +mod in_subquery; pub mod utils;