diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 830f8881645..1c11bd8de6d 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -540,6 +540,7 @@ jobs: org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometPartitionAggregateWindowSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometSmjJoinFilterFuzzSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 20f815854e6..1cf2f06d213 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -188,6 +188,7 @@ jobs: org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometPartitionAggregateWindowSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometSmjJoinFilterFuzzSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/docs/source/user-guide/latest/tuning.md b/docs/source/user-guide/latest/tuning.md index 13fabf00bb5..e29e45477cb 100644 --- a/docs/source/user-guide/latest/tuning.md +++ b/docs/source/user-guide/latest/tuning.md @@ -602,11 +602,9 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n Spark one. - `sort` over every leaf of the sorted rows (Comet also 0.15 ns per byte beyond 12 per leaf); `sortSpill` adds the price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj` - the probe side of a broadcast hash join, over every output leaf. A sort-merge join with a join condition adds - `smjCondition`, also over every output leaf: Comet builds every pair of rows of equal keys before the condition drops - them, about 3.5 times Spark's price on band joins. A condition that is one validity interval, `L <= V < U` with `V` - from one input and `L` and `U` from the other (casts, date truncations, `COALESCE(U, literal)` and `U IS NULL OR` - allowed), adds nothing when `U` is not `L` shifted by a constant through the projections below the join. + the probe side of a broadcast hash join, over every output leaf. A sort-merge join condition holding a `CASE` or an + `IF` over columns of both inputs adds `smjCrossCondition`, also over every output leaf: natively such a condition + cost 1.7 to 4.5 times Spark's price per pair of rows of equal keys. Any other condition adds nothing. - `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated, diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 3e07bdda05b..bb85467b9d6 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -442,3 +442,11 @@ harness = false [[bench]] name = "nested_comparison" harness = false + +[[bench]] +name = "smj_interval_filter" +harness = false + +[[bench]] +name = "smj_streamed_filter" +harness = false diff --git a/native/spark-expr/benches/smj_interval_filter.rs b/native/spark-expr/benches/smj_interval_filter.rs new file mode 100644 index 00000000000..85062a45f0a --- /dev/null +++ b/native/spark-expr/benches/smj_interval_filter.rs @@ -0,0 +1,262 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Sort-merge join with a validity-interval join filter, shaped like a Spark +//! `o.t > cast(r.eff as timestamp) AND o.t <= cast(r.next_eff as timestamp)` +//! lookup: a string key, large buffered key groups spread over several input +//! batches, a wide streamed side, and Spark's date-to-timestamp casts. + +use arrow::array::{ + ArrayRef, Date32Array, Decimal128Array, Int64Array, RecordBatch, StringArray, + TimestampMicrosecondArray, +}; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}; +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; +use datafusion::common::{JoinSide, JoinType, NullEquality}; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::execution::TaskContext; +use datafusion::logical_expr::Operator; +use datafusion::physical_expr::expressions::{BinaryExpr, Column}; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_plan::joins::utils::{ColumnIndex, JoinFilter}; +use datafusion::physical_plan::joins::SortMergeJoinExec; +use datafusion::physical_plan::{collect, ExecutionPlan}; +use datafusion::prelude::SessionConfig; +use datafusion_comet_spark_expr::{Cast, EvalMode, SparkCastOptions}; +use std::sync::Arc; +use tokio::runtime::Runtime; + +const MICROS_PER_DAY: i64 = 86_400_000_000; + +fn timestamp_type() -> DataType { + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())) +} + +fn key_of(k: usize) -> String { + format!("C{k:06}") +} + +/// Streamed side: `key`, the lookup time `t` and `payload_cols` payload +/// columns alternating Int64 and Utf8. Every row falls into one interval of +/// its key group. +fn build_streamed( + keys: usize, + rows_per_key: usize, + group_rows: usize, + payload_cols: usize, + batch_size: usize, +) -> (SchemaRef, Vec) { + let mut fields = vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, true)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let mut rows: Vec<(usize, i64)> = (0..num_rows) + .map(|i| { + let day = (i * 7919) % group_rows; + let offset = ((i * 104_729) as i64 % MICROS_PER_DAY) + 1; + (i / rows_per_key, day as i64 * MICROS_PER_DAY + offset) + }) + .collect(); + rows.sort(); + let mut columns: Vec = vec![ + Arc::new(StringArray::from_iter_values( + rows.iter().map(|(k, _)| key_of(*k)), + )), + Arc::new( + TimestampMicrosecondArray::from_iter_values(rows.iter().map(|(_, t)| *t)) + .with_timezone("UTC"), + ), + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:08}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +/// Buffered side: `group_rows` consecutive one-day intervals +/// `(eff, next_eff]` per key, as dates. +fn build_buffered( + keys: usize, + group_rows: usize, + batch_size: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Utf8, false), + Field::new("eff", DataType::Date32, true), + Field::new("next_eff", DataType::Date32, true), + Field::new("rate", DataType::Decimal128(18, 6), true), + ])); + let num_rows = keys * group_rows; + let key = StringArray::from_iter_values((0..num_rows).map(|i| key_of(i / group_rows))); + let eff = Date32Array::from_iter_values((0..num_rows).map(|i| (i % group_rows) as i32)); + let next_eff = + Date32Array::from_iter_values((0..num_rows).map(|i| (i % group_rows) as i32 + 1)); + let rate = Decimal128Array::from_iter_values((0..num_rows).map(|i| i as i128)) + .with_precision_and_scale(18, 6) + .unwrap(); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(key), + Arc::new(eff), + Arc::new(next_eff), + Arc::new(rate), + ], + ) + .unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +fn to_timestamp(column: &str, index: usize) -> Arc { + Arc::new(Cast::new( + Arc::new(Column::new(column, index)), + timestamp_type(), + SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + None, + None, + )) +} + +/// `t > cast(eff as timestamp) AND t <= cast(next_eff as timestamp)` +fn interval_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let t: Arc = Arc::new(Column::new("t", 0)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + Arc::clone(&t), + Operator::Gt, + to_timestamp("eff", 1), + )), + Operator::And, + Arc::new(BinaryExpr::new( + t, + Operator::LtEq, + to_timestamp("next_eff", 2), + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ColumnIndex { + index: 2, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(1).clone(), + buffered.field(1).clone(), + buffered.field(2).clone(), + ])), + ) +} + +fn make_exec(batches: &[RecordBatch], schema: &SchemaRef) -> Arc { + MemorySourceConfig::try_new_exec(&[batches.to_vec()], Arc::clone(schema), None).unwrap() +} + +fn bench_smj_interval_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("smj_interval_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows, input_batch_size) in [ + ("group11k_b8192", 4, 100, 11_000, 8192), + ("group11k_b1024", 4, 100, 11_000, 1024), + ("unique", 100_000, 1, 1, 8192), + ] { + let (streamed_schema, streamed_batches) = + build_streamed(keys, rows_per_key, group_rows, 98, input_batch_size); + let (buffered_schema, buffered_batches) = + build_buffered(keys, group_rows, input_batch_size); + let pairs = keys * rows_per_key * group_rows; + for join_type in [JoinType::Left, JoinType::Full, JoinType::Inner] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on = vec![( + Arc::new(Column::new("key", 0)) as Arc, + Arc::new(Column::new("key", 0)) as Arc, + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(interval_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = Arc::new( + TaskContext::default() + .with_session_config(SessionConfig::new().with_batch_size(8192)), + ); + let rows: usize = rt.block_on(async { + let batches = collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if join_type != JoinType::Full { + assert_eq!(rows, keys * rows_per_key); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj_interval_filter); +criterion_main!(benches); diff --git a/native/spark-expr/benches/smj_streamed_filter.rs b/native/spark-expr/benches/smj_streamed_filter.rs new file mode 100644 index 00000000000..9d6455faca6 --- /dev/null +++ b/native/spark-expr/benches/smj_streamed_filter.rs @@ -0,0 +1,297 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Sort-merge join whose filter casts a streamed column, shaped like a Spark +//! `cast(o.completed as timestamp) < b.t AND o.t > b.t` lookup of earlier +//! events: a string key, large key groups on both sides, a wide streamed +//! side and a narrow buffered one. `completed` is either a +//! `yyyy-MM-dd HH:mm:ss` string or a date. + +use arrow::array::{ + ArrayRef, Date32Array, Int64Array, RecordBatch, StringArray, TimestampMicrosecondArray, +}; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}; +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; +use datafusion::common::{JoinSide, JoinType, NullEquality}; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::execution::TaskContext; +use datafusion::logical_expr::Operator; +use datafusion::physical_expr::expressions::{BinaryExpr, Column}; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_plan::joins::utils::{ColumnIndex, JoinFilter}; +use datafusion::physical_plan::joins::SortMergeJoinExec; +use datafusion::physical_plan::{collect, ExecutionPlan}; +use datafusion::prelude::SessionConfig; +use datafusion_comet_spark_expr::{Cast, EvalMode, SparkCastOptions}; +use std::sync::Arc; +use tokio::runtime::Runtime; + +const MICROS_PER_DAY: i64 = 86_400_000_000; +const BASE_MICROS: i64 = 20_000 * MICROS_PER_DAY; +const BATCH_SIZE: usize = 8192; + +fn timestamp_type() -> DataType { + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())) +} + +fn key_of(k: usize) -> String { + format!("V{k:08}") +} + +#[derive(Clone, Copy)] +enum Completed { + String, + Date, +} + +/// Streamed side: `key`, the event time `t`, `completed` and +/// `payload_cols` payload columns alternating Int64 and Utf8. +/// +/// With `all_pass`, `completed` precedes and `t` follows every buffered +/// time of the key; otherwise `completed` lies `window_days` before `t`. +fn build_streamed( + keys: usize, + rows_per_key: usize, + group_days: i64, + window_days: i64, + all_pass: bool, + completed: Completed, + payload_cols: usize, +) -> (SchemaRef, Vec) { + let completed_type = match completed { + Completed::String => DataType::Utf8, + Completed::Date => DataType::Date32, + }; + let mut fields = vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + Field::new("completed", completed_type, true), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, true)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let rows: Vec<(usize, i64, i64)> = (0..num_rows) + .map(|i| { + if all_pass { + let t = BASE_MICROS + (group_days + 1) * MICROS_PER_DAY + i as i64; + (i / rows_per_key, t, BASE_MICROS - MICROS_PER_DAY) + } else { + let day = window_days + (i as i64 * 7919) % (group_days - window_days); + let t = BASE_MICROS + day * MICROS_PER_DAY + (i as i64 * 104_729) % MICROS_PER_DAY; + (i / rows_per_key, t, t - window_days * MICROS_PER_DAY) + } + }) + .collect(); + let completed_column: ArrayRef = match completed { + Completed::String => Arc::new(StringArray::from_iter_values(rows.iter().map(|r| { + chrono::DateTime::from_timestamp_micros(r.2) + .unwrap() + .format("%Y-%m-%d %H:%M:%S") + .to_string() + }))), + Completed::Date => Arc::new(Date32Array::from_iter_values( + rows.iter().map(|r| r.2.div_euclid(MICROS_PER_DAY) as i32), + )), + }; + let mut columns: Vec = vec![ + Arc::new(StringArray::from_iter_values( + rows.iter().map(|r| key_of(r.0)), + )), + Arc::new( + TimestampMicrosecondArray::from_iter_values(rows.iter().map(|r| r.1)) + .with_timezone("UTC"), + ), + completed_column, + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:08}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, BATCH_SIZE)) +} + +/// Buffered side: `key`, `t` spread over `group_days` days and `qty`. +fn build_buffered( + keys: usize, + group_rows: usize, + group_days: i64, + batch_size: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + Field::new("qty", DataType::Int64, true), + ])); + let num_rows = keys * group_rows; + let span = group_days * MICROS_PER_DAY; + let key = StringArray::from_iter_values((0..num_rows).map(|i| key_of(i / group_rows))); + let t = TimestampMicrosecondArray::from_iter_values( + (0..num_rows).map(|i| BASE_MICROS + (i % group_rows) as i64 * span / group_rows as i64), + ) + .with_timezone("UTC"); + let qty = Int64Array::from_iter_values((0..num_rows).map(|i| i as i64)); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(key), Arc::new(t), Arc::new(qty)], + ) + .unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +/// `cast(completed as timestamp) < b.t AND o.t > b.t` +fn streamed_cast_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let completed = Arc::new(Cast::new( + Arc::new(Column::new("completed", 0)), + timestamp_type(), + SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + None, + None, + )); + let buffered_t: Arc = Arc::new(Column::new("t", 2)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + completed, + Operator::Lt, + Arc::clone(&buffered_t), + )), + Operator::And, + Arc::new(BinaryExpr::new( + Arc::new(Column::new("t", 1)), + Operator::Gt, + buffered_t, + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 2, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(2).clone(), + streamed.field(1).clone(), + buffered.field(1).clone(), + ])), + ) +} + +fn make_exec(batches: &[RecordBatch], schema: &SchemaRef) -> Arc { + MemorySourceConfig::try_new_exec(&[batches.to_vec()], Arc::clone(schema), None).unwrap() +} + +fn bench_smj_streamed_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("smj_streamed_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows, all_pass, completed) in [ + ("few_string", 4, 1000, 2000, false, Completed::String), + ("few_date", 4, 1000, 2000, false, Completed::Date), + ("all_string", 4, 200, 2000, true, Completed::String), + ] { + let group_days = 400; + let window_days = 4; + let (streamed_schema, streamed_batches) = build_streamed( + keys, + rows_per_key, + group_days, + window_days, + all_pass, + completed, + 16, + ); + let (buffered_schema, buffered_batches) = + build_buffered(keys, group_rows, group_days, BATCH_SIZE); + let pairs = keys * rows_per_key * group_rows; + for join_type in [JoinType::Left, JoinType::Full, JoinType::Inner] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on = vec![( + Arc::new(Column::new("key", 0)) as Arc, + Arc::new(Column::new("key", 0)) as Arc, + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(streamed_cast_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = + Arc::new(TaskContext::default().with_session_config( + SessionConfig::new().with_batch_size(BATCH_SIZE), + )); + let rows: usize = rt.block_on(async { + let batches = collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if all_pass { + assert_eq!(rows, pairs); + } else { + assert!(rows > keys * rows_per_key && rows < pairs / 20); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj_streamed_filter); +criterion_main!(benches); diff --git a/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs b/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs index 82610b2a54c..8b0c3a1d72c 100644 --- a/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs +++ b/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs @@ -23,14 +23,16 @@ use std::sync::Arc; -use arrow::array::{Int64Array, RecordBatch, StringArray}; +use arrow::array::{ArrayRef, Int64Array, RecordBatch, StringArray}; use arrow::compute::SortOptions; use arrow::datatypes::{DataType, Field, Schema, SchemaRef}; use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; -use datafusion_common::NullEquality; +use datafusion_common::{JoinSide, NullEquality}; use datafusion_execution::TaskContext; -use datafusion_physical_expr::expressions::col; +use datafusion_expr::Operator; +use datafusion_physical_expr::expressions::{BinaryExpr, Column, col}; use datafusion_physical_plan::collect; +use datafusion_physical_plan::joins::utils::{ColumnIndex, JoinFilter}; use datafusion_physical_plan::joins::{SortMergeJoinExec, utils::JoinOn}; use datafusion_physical_plan::test::TestMemoryExec; use tokio::runtime::Runtime; @@ -200,5 +202,189 @@ fn bench_smj(c: &mut Criterion) { group.finish(); } -criterion_group!(benches, bench_smj); +/// Streamed side of the validity-interval join: `key`, the lookup time `t` +/// and `payload_cols` payload columns alternating Int64 and Utf8. +fn build_interval_streamed( + keys: usize, + rows_per_key: usize, + group_rows: usize, + payload_cols: usize, +) -> (SchemaRef, Vec) { + let mut fields = vec![ + Field::new("key", DataType::Int64, false), + Field::new("t", DataType::Int64, false), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, false)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let key: Vec = (0..num_rows).map(|i| (i / rows_per_key) as i64).collect(); + let t: Vec = (0..num_rows) + .map(|i| ((i * 7919) % (group_rows * 10)) as i64 + 1) + .collect(); + let mut columns: Vec = vec![ + Arc::new(Int64Array::from(key)), + Arc::new(Int64Array::from(t)), + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:012}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, 8192)) +} + +/// Buffered side of the validity-interval join: `group_rows` consecutive +/// intervals `(eff, next_eff]` per key. +fn build_interval_buffered( + keys: usize, + group_rows: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Int64, false), + Field::new("eff", DataType::Int64, false), + Field::new("next_eff", DataType::Int64, false), + Field::new("rate", DataType::Utf8, false), + ])); + let num_rows = keys * group_rows; + let key: Vec = (0..num_rows).map(|i| (i / group_rows) as i64).collect(); + let eff: Vec = (0..num_rows) + .map(|i| ((i % group_rows) * 10) as i64) + .collect(); + let next_eff: Vec = eff.iter().map(|e| e + 10).collect(); + let rate = StringArray::from_iter_values((0..num_rows).map(|i| format!("rate_{i}"))); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int64Array::from(key)), + Arc::new(Int64Array::from(eff)), + Arc::new(Int64Array::from(next_eff)), + Arc::new(rate), + ], + ) + .unwrap(); + (schema, split_batch(&batch, 8192)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +/// `streamed.t > buffered.eff AND streamed.t <= buffered.next_eff` +fn interval_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let t = Arc::new(Column::new("t", 0)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + Arc::clone(&t) as _, + Operator::Gt, + Arc::new(Column::new("eff", 1)), + )), + Operator::And, + Arc::new(BinaryExpr::new( + t, + Operator::LtEq, + Arc::new(Column::new("next_eff", 2)), + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ColumnIndex { + index: 2, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(1).clone(), + buffered.field(1).clone(), + buffered.field(2).clone(), + ])), + ) +} + +/// Validity-interval join where every streamed row passes the filter for one +/// buffered row of its key group: over large key groups, where almost every +/// pair fails, and over unique keys, where every pair passes. +fn bench_smj_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("sort_merge_join_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows) in [ + ("interval", 4, 100, 11_000), + ("interval_unique", 100_000, 1, 1), + ] { + let (streamed_schema, streamed_batches) = + build_interval_streamed(keys, rows_per_key, group_rows, 98); + let (buffered_schema, buffered_batches) = + build_interval_buffered(keys, group_rows); + let pairs = keys * rows_per_key * group_rows; + for join_type in [ + datafusion_common::JoinType::Left, + datafusion_common::JoinType::Full, + datafusion_common::JoinType::Inner, + ] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on: JoinOn = vec![( + col("key", &streamed_schema).unwrap(), + col("key", &buffered_schema).unwrap(), + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(interval_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = Arc::new(TaskContext::default()); + let rows: usize = rt.block_on(async { + let batches = + collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if join_type != datafusion_common::JoinType::Full { + assert_eq!(rows, keys * rows_per_key); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj, bench_smj_filter); criterion_main!(benches); diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs index 4fc6cccaa88..fa86b9f2f01 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs @@ -26,15 +26,19 @@ use std::sync::Arc; use arrow::array::{ - Array, ArrayBuilder, ArrayRef, BooleanArray, BooleanBuilder, RecordBatch, - RecordBatchOptions, UInt64Array, UInt64Builder, new_null_array, + Array, ArrayBuilder, ArrayRef, BooleanArray, BooleanBufferBuilder, BooleanBuilder, + RecordBatch, RecordBatchOptions, UInt64Array, UInt64Builder, new_null_array, }; use arrow::compute::kernels::zip::zip; use arrow::compute::{self, filter_record_batch}; -use arrow::datatypes::SchemaRef; +use arrow::datatypes::{Field, Schema, SchemaRef}; +use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion}; use datafusion_common::{JoinSide, JoinType, Result}; +use datafusion_physical_expr::expressions::Column; +use datafusion_physical_expr::utils::collect_columns; +use datafusion_physical_expr_common::physical_expr::{PhysicalExprRef, is_volatile}; -use crate::joins::utils::JoinFilter; +use crate::joins::utils::{ColumnIndex, JoinFilter}; /// Metadata for tracking filter results during deferred filtering /// @@ -145,82 +149,31 @@ pub fn needs_deferred_filtering( && matches!(join_type, JoinType::Left | JoinType::Right | JoinType::Full) } -/// Gets the arrays which join filters are applied on -/// -/// Extracts the columns needed for filter evaluation from left and right batch columns -pub fn get_filter_columns( - join_filter: &Option, - left_columns: &[ArrayRef], - right_columns: &[ArrayRef], -) -> Vec { - let mut filter_columns = vec![]; - - if let Some(f) = join_filter { - let left_columns: Vec = f - .column_indices() - .iter() - .filter(|col_index| col_index.side == JoinSide::Left) - .map(|i| Arc::clone(&left_columns[i.index])) - .collect(); - let right_columns: Vec = f - .column_indices() - .iter() - .filter(|col_index| col_index.side == JoinSide::Right) - .map(|i| Arc::clone(&right_columns[i.index])) - .collect(); - - filter_columns.extend(left_columns); - filter_columns.extend(right_columns); - } - - filter_columns -} - -/// Determines if current index is the last occurrence of a row +/// Marks the entries that are the last of their input row /// /// Used during filter mask correction to detect row boundaries when grouping -/// output rows by input row. -fn last_index_for_row( - row_index: usize, - indices: &UInt64Array, - batch_ids: &[usize], - indices_len: usize, -) -> bool { - debug_assert_eq!( - indices.len(), - indices_len, - "indices.len() should match indices_len parameter" - ); +/// output rows by input row. Entries without a row index (null-joined rows +/// that belong to no input row group) are not marked and do not end a +/// group: a FULL join stages the null-joined rows of unmatched buffered rows +/// ahead of the pairs a freeze materializes, so they can fall between two +/// runs of one streamed row's pairs. +fn last_entries_of_rows(indices: &UInt64Array, batch_ids: &[usize]) -> Vec { debug_assert_eq!( batch_ids.len(), - indices_len, - "batch_ids.len() should match indices_len" - ); - debug_assert!( - row_index < indices_len, - "row_index {row_index} should be < indices_len {indices_len}", + indices.len(), + "batch_ids.len() should match indices.len()" ); - - // If this is the last index overall, it's definitely the last for this row - if row_index == indices_len - 1 { - return true; - } - - // Check if next row has different (batch_id, index) pair - let current_batch_id = batch_ids[row_index]; - let next_batch_id = batch_ids[row_index + 1]; - - if current_batch_id != next_batch_id { - return true; - } - - // Same batch_id, check if row index is different - // Both current and next should be non-null (already joined rows) - if indices.is_null(row_index) || indices.is_null(row_index + 1) { - return true; + let mut last = vec![false; indices.len()]; + let mut next_row: Option<(usize, u64)> = None; + for i in (0..indices.len()).rev() { + if indices.is_null(i) { + continue; + } + let row = (batch_ids[i], indices.value(i)); + last[i] = next_row != Some(row); + next_row = Some(row); } - - indices.value(row_index) != indices.value(row_index + 1) + last } /// Corrects the filter mask for joins with deferred filtering @@ -260,9 +213,8 @@ pub fn get_corrected_filter_mask( // discard (null) remaining matches, null-join if none passed. // Null metadata entries are already-null-joined rows that // flow through unchanged to preserve output ordering. - for i in 0..row_indices_length { - let last_index = - last_index_for_row(i, row_indices, batch_ids, row_indices_length); + let last_entries = last_entries_of_rows(row_indices, batch_ids); + for (i, &last_index) in last_entries.iter().enumerate() { if filter_mask.is_null(i) { corrected_mask.append_value(true); } else if filter_mask.value(i) { @@ -294,6 +246,61 @@ pub fn get_corrected_filter_mask( } } +/// Selects the pairs of one freeze that deferred filtering may output +/// +/// For each streamed row, `get_corrected_filter_mask` keeps every pair that +/// passed the filter, or null-joins the row's last pair when none did, and +/// discards the rest. A row's pairs can span several freezes, so a freeze +/// cannot tell which of its failing pairs ends up null-joined. Within each +/// run of one row's pairs it therefore keeps the passing pairs, or only the +/// run's last pair when none passed. Over the runs of one row this keeps +/// all of its passing pairs and its overall last pair, which is all +/// `get_corrected_filter_mask` outputs from it, and the discarded failing +/// pairs never decide which pair is null-joined. +/// +/// # Arguments +/// * `row_indices` - Which streamed row produced each pair (no nulls) +/// * `filter_mask` - Whether each pair passed the filter (no nulls) +/// +/// # Returns +/// A mask that is `true` for the pairs to materialize and keep +pub fn deferred_filter_candidates( + row_indices: &UInt64Array, + filter_mask: &BooleanArray, +) -> BooleanArray { + debug_assert_eq!( + row_indices.len(), + filter_mask.len(), + "row_indices and filter_mask must have same length" + ); + debug_assert_eq!(row_indices.null_count(), 0); + debug_assert_eq!(filter_mask.null_count(), 0); + + let rows = row_indices.values(); + let passed = filter_mask.values(); + let mut keep = BooleanBufferBuilder::new(rows.len()); + let mut run_start = 0; + for i in 0..rows.len() { + if i + 1 < rows.len() && rows[i + 1] == rows[i] { + continue; + } + let run_len = i + 1 - run_start; + if run_len == 1 { + keep.append(true); + } else { + let run = passed.slice(run_start, run_len); + if run.count_set_bits() > 0 { + keep.append_buffer(&run); + } else { + keep.append_n(run_len - 1, false); + keep.append(true); + } + } + run_start = i + 1; + } + BooleanArray::new(keep.finish(), None) +} + /// Applies corrected filter mask to record batch based on join type /// /// The corrected mask has three possible values per row: @@ -386,3 +393,199 @@ pub fn filter_record_batch_by_join_type( JoinType::Inner => Ok(filter_record_batch(record_batch, corrected_mask)?), } } + +/// A join filter whose largest subexpressions over the columns of one input +/// alone are lifted out, so they can be evaluated once per buffered or +/// streamed row instead of once per pair. +#[derive(Debug)] +pub(super) struct HoistedJoinFilter { + /// The filter over `schema`, reading each lifted subexpression's result + /// from a column + pub expression: PhysicalExprRef, + /// Schema of the batch `expression` is evaluated against + pub schema: SchemaRef, + /// Where each column of `schema` comes from + pub inputs: Vec, + /// The lifted subexpressions over buffered columns, over the buffered + /// input's columns + pub buffered_exprs: Vec, + /// The lifted subexpressions over streamed columns, over the streamed + /// input's columns in `streamed_projection` + pub streamed_exprs: Vec, + /// The streamed input's columns `streamed_exprs` read + pub streamed_projection: Vec, +} + +#[derive(Debug)] +pub(super) enum HoistedFilterInput { + /// A column of the streamed or the buffered input + Column(ColumnIndex), + /// The result of `buffered_exprs[i]` + Buffered(usize), + /// The result of `streamed_exprs[i]` + Streamed(usize), +} + +impl HoistedJoinFilter { + /// Lifts the non-volatile subexpressions of `filter` that read the + /// columns of one input and none of the other. Returns `None` when there + /// is none besides bare columns. + pub fn try_new( + filter: &JoinFilter, + buffered_side: JoinSide, + streamed_schema: &Schema, + buffered_schema: &Schema, + ) -> Result> { + let column_indices = filter.column_indices(); + let num_filter_columns = column_indices.len(); + let mut lifted: Vec<(JoinSide, PhysicalExprRef)> = vec![]; + let expression = Arc::clone(filter.expression()) + .transform_down(|expr| { + if expr.downcast_ref::().is_some() || is_volatile(&expr) { + return Ok(Transformed::no(expr)); + } + let columns = collect_columns(&expr); + let Some(side) = columns + .iter() + .next() + .map(|c| column_indices[c.index()].side) + else { + return Ok(Transformed::no(expr)); + }; + if columns + .iter() + .any(|c| column_indices[c.index()].side != side) + { + return Ok(Transformed::no(expr)); + } + let position = match lifted.iter().position(|(_, e)| **e == *expr) { + Some(position) => position, + None => { + lifted.push((side, Arc::clone(&expr))); + lifted.len() - 1 + } + }; + Ok(Transformed::new( + Arc::new(Column::new( + &lifted_name(&lifted, position, buffered_side), + num_filter_columns + position, + )), + true, + TreeNodeRecursion::Jump, + )) + })? + .data; + if lifted.is_empty() { + return Ok(None); + } + + let mut used: Vec = collect_columns(&expression) + .iter() + .map(|c| c.index()) + .filter(|&index| index < num_filter_columns) + .collect(); + used.sort_unstable(); + used.dedup(); + let mut fields: Vec = used + .iter() + .map(|&index| filter.schema().field(index).clone()) + .collect(); + let mut inputs: Vec = used + .iter() + .map(|&index| HoistedFilterInput::Column(column_indices[index].clone())) + .collect(); + let mut buffered_exprs = vec![]; + let mut streamed_exprs = vec![]; + for (position, (side, expr)) in lifted.iter().enumerate() { + fields.push(Field::new( + lifted_name(&lifted, position, buffered_side), + expr.data_type(filter.schema())?, + expr.nullable(filter.schema())?, + )); + if *side == buffered_side { + inputs.push(HoistedFilterInput::Buffered(buffered_exprs.len())); + buffered_exprs.push(Arc::clone(expr)); + } else { + inputs.push(HoistedFilterInput::Streamed(streamed_exprs.len())); + streamed_exprs.push(Arc::clone(expr)); + } + } + + let expression = expression + .transform_up(|expr| { + let Some(column) = expr.downcast_ref::() else { + return Ok(Transformed::no(expr)); + }; + let index = match used.binary_search(&column.index()) { + Ok(position) => position, + Err(_) => used.len() + column.index() - num_filter_columns, + }; + Ok(Transformed::yes( + Arc::new(Column::new(column.name(), index)) as PhysicalExprRef, + )) + })? + .data; + let rebind = |exprs: Vec, + schema: &Schema, + projection: Option<&[usize]>| { + exprs + .into_iter() + .map(|expr| { + expr.transform_up(|expr| { + let Some(column) = expr.downcast_ref::() else { + return Ok(Transformed::no(expr)); + }; + let index = column_indices[column.index()].index; + let position = match projection { + Some(projection) => projection.binary_search(&index).unwrap(), + None => index, + }; + Ok(Transformed::yes(Arc::new(Column::new( + schema.field(index).name(), + position, + )) + as PhysicalExprRef)) + }) + .map(|transformed| transformed.data) + }) + .collect::>>() + }; + let mut streamed_projection: Vec = streamed_exprs + .iter() + .flat_map(collect_columns) + .map(|c| column_indices[c.index()].index) + .collect(); + streamed_projection.sort_unstable(); + streamed_projection.dedup(); + let buffered_exprs = rebind(buffered_exprs, buffered_schema, None)?; + let streamed_exprs = + rebind(streamed_exprs, streamed_schema, Some(&streamed_projection))?; + + Ok(Some(Self { + expression, + schema: Arc::new(Schema::new(fields)), + inputs, + buffered_exprs, + streamed_exprs, + streamed_projection, + })) + } +} + +/// Name of the column holding the result of `lifted[position]` +fn lifted_name( + lifted: &[(JoinSide, PhysicalExprRef)], + position: usize, + buffered_side: JoinSide, +) -> String { + let side = lifted[position].0; + let side_position = lifted[..position] + .iter() + .filter(|(s, _)| *s == side) + .count(); + if side == buffered_side { + format!("__buffered_{side_position}") + } else { + format!("__streamed_{side_position}") + } +} diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs index 43306248d03..ebd976fba43 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs @@ -30,8 +30,9 @@ use std::ops::Range; use std::sync::Arc; use crate::joins::sort_merge_join::filter::{ - FilterMetadata, filter_record_batch_by_join_type, get_corrected_filter_mask, - get_filter_columns, needs_deferred_filtering, + FilterMetadata, HoistedFilterInput, HoistedJoinFilter, deferred_filter_candidates, + filter_record_batch_by_join_type, get_corrected_filter_mask, + needs_deferred_filtering, }; use crate::joins::sort_merge_join::metrics::SortMergeJoinMetrics; use crate::joins::utils::{JoinFilter, JoinKeyComparator}; @@ -49,7 +50,7 @@ use arrow::datatypes::SchemaRef; use datafusion_common::cast::as_uint64_array; use datafusion_common::instant::Instant; use datafusion_common::{ - DataFusionError, JoinType, NullEquality, Result, exec_err, internal_err, + DataFusionError, JoinSide, JoinType, NullEquality, Result, exec_err, internal_err, }; use datafusion_execution::memory_pool::MemoryReservation; use datafusion_execution::runtime_env::RuntimeEnv; @@ -156,6 +157,50 @@ impl StreamedBatch { } self.num_output_rows += 1; } + + /// Appends the pairs of the current streamed index with each buffered + /// index in `buffered_indices` of the buffered batch with + /// `buffered_batch_idx` index. + #[inline(never)] + fn append_output_pairs( + &mut self, + buffered_batch_idx: usize, + buffered_indices: Range, + batch_size: usize, + ) { + if buffered_indices.is_empty() { + return; + } + if self.output_indices.is_empty() + || self.buffered_batch_idx != Some(buffered_batch_idx) + { + debug_assert!( + batch_size > self.num_output_rows, + "batch_size ({batch_size}) must be > num_output_rows ({})", + self.num_output_rows + ); + let capacity = batch_size - self.num_output_rows; + self.output_indices.push(StreamedJoinedChunk { + buffered_batch_idx: Some(buffered_batch_idx), + streamed_indices: UInt64Builder::with_capacity(capacity), + buffered_indices: UInt64Builder::with_capacity(capacity), + }); + self.buffered_batch_idx = Some(buffered_batch_idx); + } + let current_chunk = self.output_indices.last_mut().unwrap(); + let num_pairs = buffered_indices.len(); + current_chunk + .streamed_indices + .append_value_n(self.idx as u64, num_pairs); + let buffered_builder = &mut current_chunk.buffered_indices; + buffered_builder.append_value_n(0, num_pairs); + let values = buffered_builder.values_slice_mut(); + let appended = values.len() - num_pairs..; + for (value, idx) in values[appended].iter_mut().zip(buffered_indices) { + *value = idx as u64; + } + self.num_output_rows += num_pairs; + } } /// Per-row filter outcome tracking for full outer joins. @@ -166,7 +211,7 @@ impl StreamedBatch { /// cannot distinguish "never matched" (handled by [`BufferedBatch::null_joined`]) /// from "matched but all filters failed" (must be emitted as null-joined). #[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] pub(super) enum FilterState { /// Row never appeared in a matched pair. Unvisited = 0, @@ -210,6 +255,18 @@ pub(super) struct BufferedBatch { /// but if batch is spilled to disk this property is preferable /// and less expensive pub num_rows: usize, + /// Results of the hoisted join filter subexpressions over a range of + /// this batch, tracked in `reserved_amount` + pub filter_cache: Option, +} + +/// Results of [`HoistedJoinFilter::buffered_exprs`] over the rows `range` +/// of a buffered batch. +#[derive(Debug)] +pub(super) struct BufferedFilterCache { + range: Range, + columns: Vec, + mem: usize, } impl BufferedBatch { @@ -248,6 +305,7 @@ impl BufferedBatch { reserved_amount: 0, join_filter_status: vec![FilterState::Unvisited; num_rows], num_rows, + filter_cache: None, }) } } @@ -290,6 +348,9 @@ pub(super) struct MaterializingSortMergeJoinStream { pub sort_options: Vec, /// optional join filter pub filter: Option, + /// `filter` with its buffered-only subexpressions lifted out, when it + /// has any + pub hoisted_filter: Option, /// How the join is performed pub join_type: JoinType, /// Cached `needs_deferred_filtering(filter, join_type)` — both inputs @@ -551,6 +612,23 @@ impl MaterializingSortMergeJoinStream { semi/anti/mark joins use BitwiseSortMergeJoinStream" ); let join_time = join_metrics.join_time(); + let buffered_side = if join_type == JoinType::Right { + JoinSide::Left + } else { + JoinSide::Right + }; + let hoisted_filter = filter + .as_ref() + .map(|filter| { + HoistedJoinFilter::try_new( + filter, + buffered_side, + &streamed_schema, + &buffered_schema, + ) + }) + .transpose()? + .flatten(); let mut this = Self { sort_options, null_equality, @@ -568,6 +646,7 @@ impl MaterializingSortMergeJoinStream { on_buffered, deferred_filtering: needs_deferred_filtering(&filter, join_type), filter, + hoisted_filter, joined_record_batches: JoinedRecordBatches { joined_batches: new_output_coalescer(Arc::clone(&schema), batch_size), filter_metadata: FilterMetadata::new(), @@ -694,13 +773,24 @@ impl MaterializingSortMergeJoinStream { while !self.buffered_data.scanning_finished() && self.num_unfrozen_pairs() < self.batch_size { - let scanning_idx = self.buffered_data.scanning_idx(); - self.streamed_batch.append_output_pair( - Some(self.buffered_data.scanning_batch_idx), - Some(scanning_idx), - self.batch_size, - ); - self.buffered_data.scanning_advance(); + let range = &self.buffered_data.scanning_batch().range; + let scanning_idx = range.start + self.buffered_data.scanning_offset; + let num_pairs = (range.end - scanning_idx) + .min(self.batch_size - self.num_unfrozen_pairs()); + if num_pairs == 1 { + self.streamed_batch.append_output_pair( + Some(self.buffered_data.scanning_batch_idx), + Some(scanning_idx), + self.batch_size, + ); + } else { + self.streamed_batch.append_output_pairs( + self.buffered_data.scanning_batch_idx, + scanning_idx..scanning_idx + num_pairs, + self.batch_size, + ); + } + self.buffered_data.scanning_advance_by(num_pairs); } if self.num_unfrozen_pairs() >= self.batch_size { return false; @@ -1504,10 +1594,15 @@ impl MaterializingSortMergeJoinStream { Ok(()) } - /// Materializes columns, evaluates the join filter, and pushes output + /// Evaluates the join filter, materializes columns, and pushes output /// for all matched chunks in a single batch. This avoids per-chunk /// RecordBatch construction and filter evaluation, which dominates /// cost when keys are near-unique (1 row per chunk). + /// + /// The filter is evaluated on the filter columns alone; the remaining + /// columns are materialized only for the pairs that can reach the output, + /// so a large key group whose pairs mostly fail the filter does not + /// gather every column for every pair. fn freeze_streamed_matched( &mut self, matched_chunks: &[(usize, UInt64Array, UInt64Array)], @@ -1541,98 +1636,499 @@ impl MaterializingSortMergeJoinStream { as_uint64_array(&compute::concat(&refs)?)?.clone() }; - let left_columns = - materialize_left_columns(&self.streamed_batch.batch, &combined_left_indices)?; + let hoisted = self + .hoisted_filter + .as_ref() + .is_some_and(|hoisted| hoisted.buffered_exprs.is_empty()) + || self.cache_hoisted_filter(matched_chunks)?; + let evaluation = if hoisted { + Some(self.evaluate_hoisted_join_filter( + &combined_left_indices, + matched_chunks, + total_matched_rows, + )?) + } else { + self.evaluate_join_filter( + &combined_left_indices, + matched_chunks, + total_matched_rows, + )? + }; + let Some(evaluation) = evaluation else { + let output_batch = self.materialize_output_batch( + &combined_left_indices, + matched_chunks, + total_matched_rows, + None, + )?; + self.joined_record_batches + .push_batch_without_metadata(output_batch); + return Ok(()); + }; + let mask = &evaluation.mask; - let right_columns = - self.materialize_right_columns(matched_chunks, total_matched_rows)?; + // Track which buffered rows had all filter matches fail, + // so full join can emit them as null-joined later. + if self.join_type == JoinType::Full { + let mut offset = 0usize; + for (batch_idx, _left, right) in matched_chunks { + let chunk_len = right.len(); + let status = + &mut self.buffered_data.batches[*batch_idx].join_filter_status; + let chunk_mask = mask.values().slice(offset, chunk_len); + if right.null_count() == 0 { + // A row's state only rises: every pair makes it at least + // AllFailed, and a passing one SomePassed. + if let Some(range) = is_contiguous_range(right) { + for state in &mut status[range] { + *state = (*state).max(FilterState::AllFailed); + } + } else { + for &idx in right.values() { + let state = &mut status[idx as usize]; + *state = (*state).max(FilterState::AllFailed); + } + } + for i in chunk_mask.set_indices() { + status[right.value(i) as usize] = FilterState::SomePassed; + } + } else { + for (idx, passed) in right.iter().zip(chunk_mask.iter()) { + if let Some(idx) = idx { + let state = &mut status[idx as usize]; + *state = (*state).max(if passed { + FilterState::SomePassed + } else { + FilterState::AllFailed + }); + } + } + } + offset += chunk_len; + } + debug_assert_eq!( + offset, total_matched_rows, + "offset must advance through every chunk exactly once" + ); + } - let filter_columns = if self.join_type == JoinType::Right { - get_filter_columns(&self.filter, &right_columns, &left_columns) + // Deferred filtering outputs a subset of the pairs that + // `deferred_filter_candidates` keeps; inner joins output the pairs + // that passed. + let keep = if self.deferred_filtering { + deferred_filter_candidates(&combined_left_indices, mask) } else { - get_filter_columns(&self.filter, &left_columns, &right_columns) + mask.clone() }; + let num_kept = keep.true_count(); + + if num_kept == total_matched_rows { + let output_batch = self.materialize_output_batch( + &combined_left_indices, + matched_chunks, + total_matched_rows, + Some(&evaluation), + )?; + return self.push_filtered_batch(output_batch, &combined_left_indices, mask); + } + if num_kept == 0 { + return Ok(()); + } - let columns = if self.join_type != JoinType::Right { - [left_columns, right_columns].concat() + let kept_left_indices = + as_uint64_array(&compute::filter(&combined_left_indices, &keep)?)?.clone(); + let mut kept_chunks = Vec::with_capacity(matched_chunks.len()); + let mut offset = 0usize; + for (batch_idx, left, right) in matched_chunks { + let chunk_keep = keep.slice(offset, left.len()); + offset += left.len(); + if chunk_keep.true_count() == 0 { + continue; + } + kept_chunks.push(( + *batch_idx, + as_uint64_array(&compute::filter(left, &chunk_keep)?)?.clone(), + as_uint64_array(&compute::filter(right, &chunk_keep)?)?.clone(), + )); + } + let output_batch = self.materialize_output_batch( + &kept_left_indices, + &kept_chunks, + num_kept, + None, + )?; + let kept_mask = compute::filter(mask, &keep)?; + self.push_filtered_batch(output_batch, &kept_left_indices, kept_mask.as_boolean()) + } + + /// Pushes the pairs of a filtered freeze: with their filter metadata for + /// deferred filtering, or only those that passed for inner joins. + fn push_filtered_batch( + &mut self, + output_batch: RecordBatch, + left_indices: &UInt64Array, + mask: &BooleanArray, + ) -> Result<()> { + if self.deferred_filtering { + self.joined_record_batches.push_batch_with_filter_metadata( + output_batch, + left_indices, + mask, + self.streamed_batch_counter, + self.join_type, + ); + } else if mask.false_count() == 0 { + self.joined_record_batches + .push_batch_without_metadata(output_batch); } else { - [right_columns, left_columns].concat() + let filtered_batch = filter_record_batch(&output_batch, mask)?; + self.joined_record_batches + .push_batch_without_metadata(filtered_batch); + } + Ok(()) + } + + /// Evaluates the join filter over the given pairs on an intermediate + /// batch holding only the filter columns. Returns `None` when the join + /// has no filter columns to evaluate. + fn evaluate_join_filter( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + ) -> Result> { + let Some(filter) = &self.filter else { + return Ok(None); + }; + let (streamed_side, buffered_side) = if self.join_type == JoinType::Right { + (JoinSide::Right, JoinSide::Left) + } else { + (JoinSide::Left, JoinSide::Right) }; - let output_batch = RecordBatch::try_new(Arc::clone(&self.schema), columns)?; - - if !filter_columns.is_empty() { - if let Some(f) = &self.filter { - let filter_batch = - RecordBatch::try_new(Arc::clone(f.schema()), filter_columns)?; - let filter_result = f - .expression() - .evaluate(&filter_batch)? - .into_array(filter_batch.num_rows())?; - - let filter_result_mask = - datafusion_common::cast::as_boolean_array(&filter_result)?; - - // Convert NULL filter results to false — NULL means "not satisfied" - // per SQL semantics, same as Left/Right outer joins. - let mask = if filter_result_mask.null_count() > 0 { - compute::prep_null_mask_filter(filter_result_mask) + let side_projection = |side: JoinSide| { + let mut projection: Vec = filter + .column_indices() + .iter() + .filter(|col_index| col_index.side == side) + .map(|col_index| col_index.index) + .collect(); + projection.sort_unstable(); + projection.dedup(); + projection + }; + let streamed_projection = side_projection(streamed_side); + let buffered_projection = side_projection(buffered_side); + if streamed_projection.is_empty() && buffered_projection.is_empty() { + return Ok(None); + } + + let streamed_columns = if streamed_projection.is_empty() { + vec![] + } else { + materialize_left_columns( + &self.streamed_batch.batch.project(&streamed_projection)?, + left_indices, + )? + }; + let buffered_columns = if buffered_projection.is_empty() { + vec![] + } else { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + Some(&buffered_projection), + )? + }; + + // Left-side filter columns first, then right-side ones. + let filter_columns = [JoinSide::Left, JoinSide::Right] + .into_iter() + .flat_map(|side| { + filter + .column_indices() + .iter() + .filter(move |col_index| col_index.side == side) + }) + .map(|col_index| { + let (projection, columns) = if col_index.side == streamed_side { + (&streamed_projection, &streamed_columns) } else { - filter_result_mask.clone() + (&buffered_projection, &buffered_columns) }; + let pos = projection.binary_search(&col_index.index).unwrap(); + Arc::clone(&columns[pos]) + }) + .collect::>(); - if self.deferred_filtering { - self.joined_record_batches.push_batch_with_filter_metadata( - output_batch, - &combined_left_indices, - &mask, - self.streamed_batch_counter, - self.join_type, - ); - } else { - let filtered_batch = filter_record_batch(&output_batch, &mask)?; - self.joined_record_batches - .push_batch_without_metadata(filtered_batch); - } + let filter_batch = + RecordBatch::try_new(Arc::clone(filter.schema()), filter_columns)?; + let mask = evaluate_filter_mask(filter.expression(), &filter_batch)?; + + Ok(Some(FilterEvaluation { + mask, + streamed_columns: streamed_projection + .into_iter() + .zip(streamed_columns) + .collect(), + buffered_columns: buffered_projection + .into_iter() + .zip(buffered_columns) + .collect(), + })) + } + + /// Makes sure every buffered row the given pairs reference has its + /// hoisted filter subexpressions cached, computing them over the current + /// key group's rows of a batch the first time. Returns false when the + /// join has no hoisted filter or a pair falls outside the rows that can + /// be cached: a null-joined pair, or one from an earlier key group. + fn cache_hoisted_filter( + &mut self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + ) -> Result { + let Some(hoisted) = &self.hoisted_filter else { + return Ok(false); + }; + for (batch_idx, _, right) in matched_chunks { + if right.null_count() > 0 { + return Ok(false); + } + let (min, max) = right + .values() + .iter() + .fold((u64::MAX, 0u64), |(lo, hi), &index| { + (lo.min(index), hi.max(index)) + }); + let covers = |range: &Range| { + range.start as u64 <= min && max < range.end as u64 + }; + let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; + if buffered_batch + .filter_cache + .as_ref() + .is_some_and(|cache| covers(&cache.range)) + { + continue; + } + if !covers(&buffered_batch.range) { + return Ok(false); + } + let BufferedBatchState::InMemory(batch) = &buffered_batch.batch else { + return internal_err!( + "Buffered batch should have been unspilled before evaluating the join filter" + ); + }; + let range = buffered_batch.range.clone(); + let rows = batch.slice(range.start, range.len()); + let columns = hoisted + .buffered_exprs + .iter() + .map(|expr| expr.evaluate(&rows)?.into_array(rows.num_rows())) + .collect::>>()?; + let mem = columns.iter().map(|c| c.get_array_memory_size()).sum(); + if let Some(stale) = buffered_batch.filter_cache.take() { + self.reservation.shrink(stale.mem); + buffered_batch.reserved_amount -= stale.mem; + } + self.reservation.grow(mem); + buffered_batch.reserved_amount += mem; + buffered_batch.filter_cache = Some(BufferedFilterCache { + range, + columns, + mem, + }); + self.join_metrics + .peak_mem_used() + .set_max(self.reservation.size()); + } + Ok(true) + } - // Track which buffered rows had all filter matches fail, - // so full join can emit them as null-joined later. - if self.join_type == JoinType::Full { - let mut offset = 0usize; - for (batch_idx, _left, right) in matched_chunks { - let chunk_len = right.len(); - let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; - - for i in 0..chunk_len { - if right.is_null(i) { - continue; - } - let idx = right.value(i) as usize; - match buffered_batch.join_filter_status[idx] { - FilterState::SomePassed => {} - _ if mask.value(offset + i) => { - buffered_batch.join_filter_status[idx] = - FilterState::SomePassed; - } - _ => { - buffered_batch.join_filter_status[idx] = - FilterState::AllFailed; - } - } - } - offset += chunk_len; + /// Evaluates the hoisted join filter over the given pairs, whose + /// buffered rows [`Self::cache_hoisted_filter`] cached when the filter + /// lifts buffered subexpressions. + fn evaluate_hoisted_join_filter( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + ) -> Result { + let hoisted = self.hoisted_filter.as_ref().unwrap(); + let streamed_side = if self.join_type == JoinType::Right { + JoinSide::Right + } else { + JoinSide::Left + }; + let side_projection = |streamed: bool| { + let mut projection: Vec = hoisted + .inputs + .iter() + .filter_map(|input| match input { + HoistedFilterInput::Column(column) + if (column.side == streamed_side) == streamed => + { + Some(column.index) } - debug_assert_eq!( - offset, total_matched_rows, - "offset must advance through every chunk exactly once" - ); + _ => None, + }) + .collect(); + projection.sort_unstable(); + projection.dedup(); + projection + }; + let streamed_projection = side_projection(true); + let buffered_projection = side_projection(false); + + let streamed_columns = if streamed_projection.is_empty() { + vec![] + } else { + materialize_left_columns( + &self.streamed_batch.batch.project(&streamed_projection)?, + left_indices, + )? + }; + let buffered_columns = if buffered_projection.is_empty() { + vec![] + } else { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + Some(&buffered_projection), + )? + }; + let cached_columns = self + .gather_cached_filter_columns(matched_chunks, hoisted.buffered_exprs.len())?; + let streamed_lifted_columns = evaluate_streamed_filter_exprs( + hoisted, + &self.streamed_batch.batch, + left_indices, + )?; + + let filter_columns = hoisted + .inputs + .iter() + .map(|input| match input { + HoistedFilterInput::Column(column) => { + let (projection, columns) = if column.side == streamed_side { + (&streamed_projection, &streamed_columns) + } else { + (&buffered_projection, &buffered_columns) + }; + let pos = projection.binary_search(&column.index).unwrap(); + Arc::clone(&columns[pos]) + } + HoistedFilterInput::Buffered(position) => { + Arc::clone(&cached_columns[*position]) + } + HoistedFilterInput::Streamed(position) => { + Arc::clone(&streamed_lifted_columns[*position]) + } + }) + .collect::>(); + let filter_batch = RecordBatch::try_new_with_options( + Arc::clone(&hoisted.schema), + filter_columns, + &RecordBatchOptions::new().with_row_count(Some(total_matched_rows)), + )?; + let mask = evaluate_filter_mask(&hoisted.expression, &filter_batch)?; + + Ok(FilterEvaluation { + mask, + streamed_columns: streamed_projection + .into_iter() + .zip(streamed_columns) + .collect(), + buffered_columns: buffered_projection + .into_iter() + .zip(buffered_columns) + .collect(), + }) + } + + /// Gathers the cached hoisted filter subexpression results of the + /// buffered rows the given pairs reference. + fn gather_cached_filter_columns( + &self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + num_columns: usize, + ) -> Result> { + if num_columns == 0 { + return Ok(vec![]); + } + let mut pieces: Vec> = + vec![Vec::with_capacity(matched_chunks.len()); num_columns]; + for (batch_idx, _, right) in matched_chunks { + let Some(cache) = &self.buffered_data.batches[*batch_idx].filter_cache else { + return internal_err!("Hoisted join filter results were not cached"); + }; + let offset = cache.range.start; + if let Some(range) = is_contiguous_range(right) { + for (column, piece) in cache.columns.iter().zip(pieces.iter_mut()) { + piece.push(column.slice(range.start - offset, range.len())); + } + } else { + let indices = UInt64Array::from_iter_values( + right.values().iter().map(|&index| index - offset as u64), + ); + for (column, piece) in cache.columns.iter().zip(pieces.iter_mut()) { + piece.push(compute::take(column, &indices, None)?); } } - } else { - self.joined_record_batches - .push_batch_without_metadata(output_batch); } + pieces + .into_iter() + .map(|piece| { + if piece.len() == 1 { + Ok(Arc::clone(&piece[0])) + } else { + let refs: Vec<&dyn Array> = + piece.iter().map(|a| a.as_ref()).collect(); + Ok(compute::concat(&refs)?) + } + }) + .collect() + } - Ok(()) + /// Materializes the output batch of the given pairs, reusing the filter + /// columns `evaluation` gathered for the same pairs. + fn materialize_output_batch( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + evaluation: Option<&FilterEvaluation>, + ) -> Result { + let left_columns = complete_columns( + self.streamed_schema.fields().len(), + evaluation.map_or(&[], |e| e.streamed_columns.as_slice()), + |projection| match projection { + None => { + materialize_left_columns(&self.streamed_batch.batch, left_indices) + } + Some(projection) => materialize_left_columns( + &self.streamed_batch.batch.project(projection)?, + left_indices, + ), + }, + )?; + let right_columns = complete_columns( + self.buffered_schema.fields().len(), + evaluation.map_or(&[], |e| e.buffered_columns.as_slice()), + |projection| { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + projection, + ) + }, + )?; + + let columns = if self.join_type != JoinType::Right { + [left_columns, right_columns].concat() + } else { + [right_columns, left_columns].concat() + }; + Ok(RecordBatch::try_new(Arc::clone(&self.schema), columns)?) } /// Materializes right-side columns across all matched chunks. @@ -1640,11 +2136,13 @@ impl MaterializingSortMergeJoinStream { /// When chunks reference a single buffered batch, indices are concatenated /// for a single fetch. When multiple batches are involved, `interleave` /// gathers columns across sources. A null-row sentinel at source index 0 - /// handles null right indices (unmatched streamed rows). + /// handles null right indices (unmatched streamed rows). `projection` + /// selects the buffered columns to materialize (all when `None`). fn materialize_right_columns( &self, matched_chunks: &[(usize, UInt64Array, UInt64Array)], total_matched_rows: usize, + projection: Option<&[usize]>, ) -> Result> { let first_batch_idx = matched_chunks[0].0; let single_source = matched_chunks.iter().all(|c| c.0 == first_batch_idx); @@ -1662,9 +2160,19 @@ impl MaterializingSortMergeJoinStream { &self.buffered_data, first_batch_idx, &combined_right_indices, + projection, ); } + // Multiple source batches each contributing runs of consecutive + // rows, as a large key group spanning several batches does: copy the + // runs instead of gathering row by row. + if matched_chunks.len() * MIN_AVG_RUN_LEN <= total_matched_rows + && let Some(columns) = self.copy_buffered_runs(matched_chunks, projection)? + { + return Ok(columns); + } + // Multiple source batches: map each buffered_batch_idx to a // contiguous source index. A null sentinel array is prepended as // source 0 only when some right index is actually null (an @@ -1769,8 +2277,6 @@ impl MaterializingSortMergeJoinStream { } } - let num_right_cols = self.buffered_schema.fields().len(); - // Read each source batch once (spilled batches require disk I/O). let source_data: Vec<&RecordBatch> = source_batches .iter() @@ -1794,10 +2300,14 @@ impl MaterializingSortMergeJoinStream { vec![] }; + let col_indices: Vec = match projection { + Some(projection) => projection.to_vec(), + None => (0..self.buffered_schema.fields().len()).collect(), + }; let mut source_arrays: Vec<&dyn Array> = Vec::with_capacity(source_data.len() + source_offset); - let mut right_columns = Vec::with_capacity(num_right_cols); - for col_idx in 0..num_right_cols { + let mut right_columns = Vec::with_capacity(col_indices.len()); + for col_idx in col_indices { source_arrays.clear(); source_arrays.extend(null_arrays.get(col_idx).map(|a| a.as_ref())); source_arrays.extend(source_data.iter().map(|d| d.column(col_idx).as_ref())); @@ -1808,6 +2318,54 @@ impl MaterializingSortMergeJoinStream { Ok(right_columns) } + /// Copies the buffered columns of the given pairs run by run when every + /// chunk is a run of consecutive rows of its batch. Returns `None` + /// otherwise. + #[inline(never)] + fn copy_buffered_runs( + &self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + projection: Option<&[usize]>, + ) -> Result>> { + let Some(ranges) = matched_chunks + .iter() + .map(|(_, _, right)| is_contiguous_range(right)) + .collect::>>() + else { + return Ok(None); + }; + let batches = matched_chunks + .iter() + .map(|(batch_idx, _, _)| { + match &self.buffered_data.batches[*batch_idx].batch { + BufferedBatchState::InMemory(batch) => Ok(batch), + BufferedBatchState::Spilled(_) => internal_err!( + "Buffered batch should have been unspilled before fetching columns" + ), + } + }) + .collect::>>()?; + let col_indices: Vec = match projection { + Some(projection) => projection.to_vec(), + None => (0..self.buffered_schema.fields().len()).collect(), + }; + col_indices + .into_iter() + .map(|col_idx| { + let slices: Vec = batches + .iter() + .zip(&ranges) + .map(|(batch, range)| { + batch.column(col_idx).slice(range.start, range.len()) + }) + .collect(); + let refs: Vec<&dyn Array> = slices.iter().map(|a| a.as_ref()).collect(); + Ok(compute::concat(&refs)?) + }) + .collect::>>() + .map(Some) + } + fn filter_joined_batch(&mut self) -> Result { // Metadata should be aligned before processing self.joined_record_batches @@ -1853,6 +2411,13 @@ impl MaterializingSortMergeJoinStream { return Ok(record_batch); } + // No pair failed the filter: the corrected mask keeps every row. + if !out_mask.has_false() { + self.joined_record_batches + .clear(&self.schema, self.batch_size); + return Ok(record_batch); + } + // Validate inputs to get_corrected_filter_mask debug_assert_eq!( out_indices.len(), @@ -1914,6 +2479,107 @@ fn materialize_left_columns( } } +/// Evaluates the lifted streamed subexpressions of `hoisted` for the +/// streamed rows `indices` of `batch`: once per run of equal indices, as a +/// streamed row's pairs form, unless the runs are short. +fn evaluate_streamed_filter_exprs( + hoisted: &HoistedJoinFilter, + batch: &RecordBatch, + indices: &UInt64Array, +) -> Result> { + if hoisted.streamed_exprs.is_empty() { + return Ok(vec![]); + } + let evaluate = |rows: &UInt64Array| { + let projected = batch.project(&hoisted.streamed_projection)?; + let rows = RecordBatch::try_new_with_options( + projected.schema(), + materialize_left_columns(&projected, rows)?, + &RecordBatchOptions::new().with_row_count(Some(rows.len())), + )?; + hoisted + .streamed_exprs + .iter() + .map(|expr| expr.evaluate(&rows)?.into_array(rows.num_rows())) + .collect::>>() + }; + let values = indices.values(); + let num_runs = values.windows(2).filter(|w| w[0] != w[1]).count() + 1; + if num_runs * 2 > values.len() { + return evaluate(indices); + } + let mut run_rows = Vec::with_capacity(num_runs); + let mut positions = Vec::with_capacity(values.len()); + for (i, &row) in values.iter().enumerate() { + if i == 0 || row != values[i - 1] { + run_rows.push(row); + } + positions.push((run_rows.len() - 1) as u32); + } + let positions = UInt32Array::from(positions); + evaluate(&UInt64Array::from(run_rows))? + .iter() + .map(|column| Ok(compute::take(column, &positions, None)?)) + .collect() +} + +/// Average length of the runs of consecutive buffered rows from which +/// copying the runs beats gathering row by row. +const MIN_AVG_RUN_LEN: usize = 16; + +/// Evaluates a join filter expression into a mask, NULL results counting +/// as not satisfied. +fn evaluate_filter_mask( + expression: &Arc, + filter_batch: &RecordBatch, +) -> Result { + let filter_result = expression + .evaluate(filter_batch)? + .into_array(filter_batch.num_rows())?; + let filter_result_mask = datafusion_common::cast::as_boolean_array(&filter_result)?; + + // Convert NULL filter results to false — NULL means "not satisfied" + // per SQL semantics, same as Left/Right outer joins. + Ok(if filter_result_mask.null_count() > 0 { + compute::prep_null_mask_filter(filter_result_mask) + } else { + filter_result_mask.clone() + }) +} + +/// Join filter result for the pairs of one freeze, with the filter columns +/// gathered for it as `(column index, array)` per side. +struct FilterEvaluation { + mask: BooleanArray, + streamed_columns: Vec<(usize, ArrayRef)>, + buffered_columns: Vec<(usize, ArrayRef)>, +} + +/// Builds `num_columns` columns from those already `gathered`, gathering +/// the rest through `gather` (all of them when its projection is `None`). +fn complete_columns( + num_columns: usize, + gathered: &[(usize, ArrayRef)], + gather: impl FnOnce(Option<&[usize]>) -> Result>, +) -> Result> { + if gathered.is_empty() { + return gather(None); + } + let mut columns: Vec> = vec![None; num_columns]; + for (index, array) in gathered { + columns[*index] = Some(Arc::clone(array)); + } + let missing: Vec = (0..num_columns) + .filter(|&index| columns[index].is_none()) + .collect(); + if !missing.is_empty() { + for (index, array) in missing.iter().zip(gather(Some(&missing))?) { + columns[*index] = Some(array); + } + } + Ok(columns.into_iter().map(Option::unwrap).collect()) +} + fn create_unmatched_columns(schema: &SchemaRef, size: usize) -> Vec { schema .fields() @@ -1934,7 +2600,7 @@ fn produce_buffered_null_batch( // Take buffered (right) columns let right_columns = - fetch_right_columns_from_batch_by_idxs(buffered_batch, buffered_indices)?; + fetch_right_columns_from_batch_by_idxs(buffered_batch, buffered_indices, None)?; // Create null streamed (left) columns let mut left_columns = streamed_schema @@ -1981,10 +2647,12 @@ fn fetch_right_columns_by_idxs( buffered_data: &BufferedData, buffered_batch_idx: usize, buffered_indices: &UInt64Array, + projection: Option<&[usize]>, ) -> Result> { fetch_right_columns_from_batch_by_idxs( &buffered_data.batches[buffered_batch_idx], buffered_indices, + projection, ) } @@ -1992,9 +2660,18 @@ fn fetch_right_columns_by_idxs( fn fetch_right_columns_from_batch_by_idxs( buffered_batch: &BufferedBatch, buffered_indices: &UInt64Array, + projection: Option<&[usize]>, ) -> Result> { match &buffered_batch.batch { BufferedBatchState::InMemory(batch) => { + let projected; + let batch = match projection { + Some(projection) => { + projected = batch.project(projection)?; + &projected + } + None => batch, + }; if let Some(range) = is_contiguous_range(buffered_indices) { Ok(batch.slice(range.start, range.len()).columns().to_vec()) } else { @@ -2043,7 +2720,11 @@ impl BufferedData { } pub fn scanning_advance(&mut self) { - self.scanning_offset += 1; + self.scanning_advance_by(1); + } + + pub fn scanning_advance_by(&mut self, n: usize) { + self.scanning_offset += n; while !self.scanning_finished() && self.scanning_batch_finished() { self.scanning_batch_idx += 1; self.scanning_offset = 0; diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index 2300059f6ee..6f6dd75b0db 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -39,7 +39,10 @@ use crate::test::exec::BarrierExec; use crate::test::{build_table_i32, build_table_i32_two_cols}; use crate::{ExecutionPlan, RecordBatchStream, common}; use crate::{ - expressions::Column, joins::sort_merge_join::filter::get_corrected_filter_mask, + expressions::Column, + joins::sort_merge_join::filter::{ + deferred_filter_candidates, get_corrected_filter_mask, + }, joins::sort_merge_join::materializing_stream::JoinedRecordBatches, }; use arrow::array::{ @@ -3182,6 +3185,113 @@ fn build_joined_record_batches() -> Result { Ok(batches) } +/// Output of `get_corrected_filter_mask` over deferred-filter metadata, as +/// `(pair, passed)` per output row: a passing pair, a pair null-joined +/// because its row found no passing pair (`false`), or a null-metadata row. +fn corrected_output( + pairs: &[usize], + rows: &[Option], + mask: &[Option], +) -> Vec<(usize, Option)> { + use arrow::array::Array; + let row_indices = UInt64Array::from(rows.to_vec()); + let filter_mask = BooleanArray::from(mask.to_vec()); + let batch_ids = vec![1; rows.len()]; + let corrected = get_corrected_filter_mask( + Left, + &row_indices, + &batch_ids, + &filter_mask, + rows.len(), + ) + .unwrap(); + (0..rows.len()) + .filter(|&i| corrected.is_valid(i)) + .map(|i| { + let passed = filter_mask.is_valid(i).then(|| corrected.value(i)); + (pairs[i], passed) + }) + .collect() +} + +/// Reducing each freeze's pairs to `deferred_filter_candidates` leaves the +/// corrected output unchanged, wherever the freeze boundaries cut the pairs +/// of one streamed row. +#[test] +fn deferred_filter_candidates_keep_corrected_output() { + let mut state = 11u64; + let mut next = move |n: u64| { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) % n + }; + for _ in 0..2000 { + // Pairs of consecutive streamed rows, with null-metadata rows + // (unmatched rows) between some of them. + let mut rows: Vec> = vec![]; + let mut mask: Vec> = vec![]; + let pass_rate = next(4); + for row in 0..1 + next(6) { + if next(4) == 0 { + rows.push(None); + mask.push(None); + } + for _ in 0..1 + next(12) { + rows.push(Some(row)); + mask.push(Some(next(4) < pass_rate)); + } + } + let pairs: Vec = (0..rows.len()).collect(); + let expected = corrected_output(&pairs, &rows, &mask); + + // Cut the pairs into freezes and reduce each freeze's matched pairs. + let mut kept_pairs = vec![]; + let mut start = 0; + while start < rows.len() { + let end = (start + 1 + next(10) as usize).min(rows.len()); + let mut segment_start = start; + while segment_start < end { + if rows[segment_start].is_none() { + kept_pairs.push(segment_start); + segment_start += 1; + continue; + } + let mut segment_end = segment_start; + while segment_end < end && rows[segment_end].is_some() { + segment_end += 1; + } + let segment = segment_start..segment_end; + let keep = deferred_filter_candidates( + &UInt64Array::from_iter_values( + rows[segment.clone()].iter().map(|r| r.unwrap()), + ), + &BooleanArray::from_iter(mask[segment.clone()].iter().copied()), + ); + kept_pairs.extend(segment.filter(|&i| keep.value(i - segment_start))); + segment_start = segment_end; + } + start = end; + } + let kept_rows: Vec<_> = kept_pairs.iter().map(|&i| rows[i]).collect(); + let kept_mask: Vec<_> = kept_pairs.iter().map(|&i| mask[i]).collect(); + let actual = corrected_output(&kept_pairs, &kept_rows, &kept_mask); + + // A null-joined row may come from any pair of its streamed row: they + // differ only in the buffered side, which is nulled. + let as_rows = |output: Vec<(usize, Option)>| { + output + .into_iter() + .map(|(pair, passed)| match passed { + Some(false) => (rows[pair], None), + _ => (rows[pair], Some(pair)), + }) + .collect::>() + }; + assert_eq!(as_rows(actual), as_rows(expected), "{rows:?} {mask:?}"); + } +} + #[tokio::test] async fn test_left_outer_join_filtered_mask() -> Result<()> { let mut joined_batches = build_joined_record_batches()?; @@ -4187,6 +4297,716 @@ async fn join_wrapped_multi_source_freeze_with_null_buffered_index() -> Result<( Ok(()) } +/// A streamed row of the validity-interval tests: key, row id, lookup time. +type IntervalStreamedRow = (i32, i32, Option); +/// A buffered row of the validity-interval tests: key, row id, `(eff, next_eff]`. +type IntervalBufferedRow = (i32, i32, i64, i64); + +const INTERVAL_PAYLOAD_COLS: usize = 24; + +fn interval_payload(col: usize, sid: i32) -> String { + format!("p{col}_{sid}") +} + +/// Streamed table `(key, sid, t, p0..)` with `INTERVAL_PAYLOAD_COLS` +/// payload columns alternating Int64 and Utf8, split into batches. +fn build_interval_streamed( + rows: &[IntervalStreamedRow], + batch_rows: usize, +) -> Arc { + use arrow::array::{ArrayRef, Int64Array, StringArray}; + let mut fields = vec![ + Field::new("key", DataType::Int32, false), + Field::new("sid", DataType::Int32, false), + Field::new("t", DataType::Int64, true), + ]; + for c in 0..INTERVAL_PAYLOAD_COLS { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, false)); + } + let schema = Arc::new(Schema::new(fields)); + let batches = rows + .chunks(batch_rows) + .map(|chunk| { + let mut columns: Vec = vec![ + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))), + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))), + Arc::new(Int64Array::from_iter(chunk.iter().map(|r| r.2))), + ]; + for c in 0..INTERVAL_PAYLOAD_COLS { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + chunk.iter().map(|r| r.1 as i64 * 1000 + c as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + chunk.iter().map(|r| interval_payload(c, r.1)), + ))); + } + } + RecordBatch::try_new(Arc::clone(&schema), columns).unwrap() + }) + .collect::>(); + TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() +} + +/// Buffered table `(bkey, bid, eff, next_eff, val)`, split into batches. +fn build_interval_buffered( + rows: &[IntervalBufferedRow], + batch_rows: usize, +) -> Arc { + use arrow::array::{Int64Array, StringArray}; + let schema = Arc::new(Schema::new(vec![ + Field::new("bkey", DataType::Int32, false), + Field::new("bid", DataType::Int32, false), + Field::new("eff", DataType::Int64, false), + Field::new("next_eff", DataType::Int64, false), + Field::new("val", DataType::Utf8, false), + ])); + let batches = rows + .chunks(batch_rows) + .map(|chunk| { + RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))), + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))), + Arc::new(Int64Array::from_iter_values(chunk.iter().map(|r| r.2))), + Arc::new(Int64Array::from_iter_values(chunk.iter().map(|r| r.3))), + Arc::new(StringArray::from_iter_values( + chunk.iter().map(|r| format!("v{}", r.1)), + )), + ], + ) + .unwrap() + }) + .collect::>(); + TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() +} + +/// Forms of the validity-interval filter `t > eff AND t <= next_eff`, all +/// with the same result. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IntervalFilter { + /// `t > eff AND t <= next_eff` + Plain, + /// `t > eff + 0 AND t <= next_eff + 0 AND next_eff - eff > 0`: with + /// subexpressions over buffered columns alone, which the join evaluates + /// once per buffered row + Buffered, + /// `t + 0 > eff AND t - 0 <= next_eff AND t + 0 >= t - 0`: with + /// subexpressions over streamed columns alone, which the join evaluates + /// once per streamed row + Streamed, + /// `t + 0 > eff + 0 AND t - 0 <= next_eff + 0 AND next_eff - eff > 0 + /// AND t + 0 >= t - 0`: with both + Both, +} + +const INTERVAL_FILTERS: [IntervalFilter; 4] = [ + IntervalFilter::Plain, + IntervalFilter::Buffered, + IntervalFilter::Streamed, + IntervalFilter::Both, +]; + +/// The validity-interval filter in the given form, with the streamed table +/// on `streamed_side`. +fn build_interval_filter( + streamed: &Schema, + buffered: &Schema, + streamed_side: JoinSide, + form: IntervalFilter, +) -> JoinFilter { + let t_field = streamed + .field_with_name("t") + .unwrap() + .clone() + .with_nullable(true); + let eff_field = buffered + .field_with_name("eff") + .unwrap() + .clone() + .with_nullable(true); + let next_eff_field = buffered + .field_with_name("next_eff") + .unwrap() + .clone() + .with_nullable(true); + let t = ColumnIndex { + index: 2, + side: streamed_side, + }; + let eff = ColumnIndex { + index: 2, + side: streamed_side.negate(), + }; + let next_eff = ColumnIndex { + index: 3, + side: streamed_side.negate(), + }; + let (column_indices, fields, t_idx, eff_idx, next_eff_idx) = + if streamed_side == JoinSide::Left { + ( + vec![t, eff, next_eff], + vec![t_field, eff_field, next_eff_field], + 0, + 1, + 2, + ) + } else { + ( + vec![eff, next_eff, t], + vec![eff_field, next_eff_field, t_field], + 2, + 0, + 1, + ) + }; + let t_col: PhysicalExprRef = Arc::new(Column::new("t", t_idx)); + let eff: PhysicalExprRef = Arc::new(Column::new("eff", eff_idx)); + let next_eff: PhysicalExprRef = Arc::new(Column::new("next_eff", next_eff_idx)); + let zero = |expr: &PhysicalExprRef, op: Operator| -> PhysicalExprRef { + Arc::new(BinaryExpr::new( + Arc::clone(expr), + op, + Arc::new(Literal::new(ScalarValue::Int64(Some(0)))), + )) + }; + let lift_buffered = matches!(form, IntervalFilter::Buffered | IntervalFilter::Both); + let lift_streamed = matches!(form, IntervalFilter::Streamed | IntervalFilter::Both); + let (lower, upper) = if lift_buffered { + (zero(&eff, Operator::Plus), zero(&next_eff, Operator::Plus)) + } else { + (Arc::clone(&eff), Arc::clone(&next_eff)) + }; + let (t_lower, t_upper) = if lift_streamed { + (zero(&t_col, Operator::Plus), zero(&t_col, Operator::Minus)) + } else { + (Arc::clone(&t_col), Arc::clone(&t_col)) + }; + let mut expression: PhysicalExprRef = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new(Arc::clone(&t_lower), Operator::Gt, lower)), + Operator::And, + Arc::new(BinaryExpr::new(Arc::clone(&t_upper), Operator::LtEq, upper)), + )); + if lift_buffered { + expression = Arc::new(BinaryExpr::new( + expression, + Operator::And, + Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new(next_eff, Operator::Minus, eff)), + Operator::Gt, + Arc::new(Literal::new(ScalarValue::Int64(Some(0)))), + )), + )); + } + if lift_streamed { + expression = Arc::new(BinaryExpr::new( + expression, + Operator::And, + Arc::new(BinaryExpr::new(t_lower, Operator::GtEq, t_upper)), + )); + } + JoinFilter::new(expression, column_indices, Arc::new(Schema::new(fields))) +} + +/// Expected `(sid, bid)` output of the validity-interval join, in the +/// streamed order the join preserves: each streamed row's passing buffered +/// rows in buffered order, or the row null-joined when none passed. FULL +/// joins append the buffered rows no streamed row passed with. +fn expected_interval_join( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], +) -> Vec<(Option, Option)> { + let mut expected = vec![]; + let mut buffered_passed = vec![false; buffered.len()]; + for &(key, sid, t) in streamed { + let mut any = false; + for (b, &(bkey, bid, eff, next_eff)) in buffered.iter().enumerate() { + if bkey == key && t.is_some_and(|t| t > eff && t <= next_eff) { + expected.push((Some(sid), Some(bid))); + buffered_passed[b] = true; + any = true; + } + } + if !any && join_type != Inner { + expected.push((Some(sid), None)); + } + } + if join_type == Full { + for (b, &(_, bid, _, _)) in buffered.iter().enumerate() { + if !buffered_passed[b] { + expected.push((None, Some(bid))); + } + } + } + expected +} + +/// Reads `(sid, bid)` from each output row, checking that every streamed +/// and buffered column belongs to that row (or is null with it). +fn interval_join_output( + batches: &[RecordBatch], + streamed_offset: usize, + buffered_offset: usize, +) -> Vec<(Option, Option)> { + use arrow::array::{Array, AsArray}; + use arrow::datatypes::{Int32Type, Int64Type}; + let mut output = vec![]; + for batch in batches { + let sid = batch + .column(streamed_offset + 1) + .as_primitive::(); + let bid = batch + .column(buffered_offset + 1) + .as_primitive::(); + for row in 0..batch.num_rows() { + let s = sid.is_valid(row).then(|| sid.value(row)); + let b = bid.is_valid(row).then(|| bid.value(row)); + for c in (0..3 + INTERVAL_PAYLOAD_COLS).filter(|&c| c != 2) { + let col = batch.column(streamed_offset + c); + assert_eq!(col.is_valid(row), s.is_some(), "column {c}"); + } + if let Some(s) = s { + for c in 0..INTERVAL_PAYLOAD_COLS { + let col = batch.column(streamed_offset + 3 + c); + if c % 2 == 0 { + assert_eq!( + col.as_primitive::().value(row), + s as i64 * 1000 + c as i64 + ); + } else { + assert_eq!( + col.as_string::().value(row), + interval_payload(c, s) + ); + } + } + } + for c in 0..5 { + assert_eq!(batch.column(buffered_offset + c).is_valid(row), b.is_some()); + } + if let Some(b) = b { + assert_eq!( + batch + .column(buffered_offset + 4) + .as_string::() + .value(row), + format!("v{b}") + ); + } + output.push((s, b)); + } + } + output +} + +/// Runs the validity-interval join with the streamed table on the outer +/// side of `join_type` (the right side for RIGHT joins) and checks it +/// against [`expected_interval_join`]: in order where the streamed order is +/// preserved, as a multiset for FULL joins. +#[expect(clippy::too_many_arguments)] +async fn check_interval_join( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], + streamed_batch_rows: usize, + buffered_batch_rows: usize, + batch_size: usize, + spill: bool, + case: &str, +) -> Result<()> { + for form in INTERVAL_FILTERS { + check_interval_join_with_filter( + join_type, + streamed, + buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + form, + &format!("{case} filter={form:?}"), + ) + .await?; + } + Ok(()) +} + +/// [`check_interval_join`] with one form of the filter. +#[expect(clippy::too_many_arguments)] +async fn check_interval_join_with_filter( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], + streamed_batch_rows: usize, + buffered_batch_rows: usize, + batch_size: usize, + spill: bool, + form: IntervalFilter, + case: &str, +) -> Result<()> { + let streamed_plan = build_interval_streamed(streamed, streamed_batch_rows); + let buffered_plan = build_interval_buffered(buffered, buffered_batch_rows); + let streamed_side = if join_type == Right { + JoinSide::Right + } else { + JoinSide::Left + }; + let filter = build_interval_filter( + &streamed_plan.schema(), + &buffered_plan.schema(), + streamed_side, + form, + ); + let (left, right) = if join_type == Right { + (buffered_plan, streamed_plan) + } else { + (streamed_plan, buffered_plan) + }; + let on = vec![( + Arc::new(Column::new_with_schema( + if join_type == Right { "bkey" } else { "key" }, + &left.schema(), + )?) as _, + Arc::new(Column::new_with_schema( + if join_type == Right { "key" } else { "bkey" }, + &right.schema(), + )?) as _, + )]; + let (streamed_offset, buffered_offset) = if join_type == Right { + (5, 0) + } else { + (0, 3 + INTERVAL_PAYLOAD_COLS) + }; + + let mut task_ctx = TaskContext::default() + .with_session_config(SessionConfig::default().with_batch_size(batch_size)); + if spill { + task_ctx = task_ctx.with_runtime( + RuntimeEnvBuilder::new() + .with_memory_limit(100, 1.0) + .with_disk_manager_builder( + DiskManagerBuilder::default() + .with_mode(DiskManagerMode::OsTmpDirectory), + ) + .build_arc()?, + ); + } + let join = join_with_filter( + left, + right, + on, + filter, + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + )?; + let batches = common::collect(join.execute(0, Arc::new(task_ctx))?).await?; + if spill && !buffered.is_empty() { + assert!( + join.metrics().unwrap().spill_count().unwrap() > 0, + "{case}: expected spilling" + ); + } + + let mut actual = interval_join_output(&batches, streamed_offset, buffered_offset); + let mut expected = expected_interval_join(join_type, streamed, buffered); + if join_type == Full { + actual.sort(); + expected.sort(); + } + assert_eq!(actual, expected, "{case}"); + Ok(()) +} + +/// Buffered key group of `rows` consecutive intervals `(10i, 10i + 10 * width]`. +fn interval_group( + key: i32, + first_bid: i32, + rows: i32, + width: i64, +) -> Vec { + (0..rows) + .map(|i| { + ( + key, + first_bid + i, + i as i64 * 10, + i as i64 * 10 + 10 * width, + ) + }) + .collect() +} + +/// One streamed row against a buffered key group much larger than the batch +/// size and spanning several buffered batches, passing the filter for none, +/// one or a few of its rows (or NULL when the streamed time is NULL). The +/// join filter is evaluated for every pair, but only the pairs that can be +/// output are materialized; the result must be the same either way. +#[tokio::test] +async fn join_filter_single_streamed_row_large_buffered_group() -> Result<()> { + let group = interval_group(1, 0, 1000, 1); + let overlapping = interval_group(1, 0, 1000, 3); + let cases: Vec<(&str, Option, &[IntervalBufferedRow])> = vec![ + ("none", Some(-5), &group), + ("null", None, &group), + ("first", Some(5), &group), + ("middle", Some(7775), &group), + ("last", Some(9999), &group), + ("several", Some(7775), &overlapping), + ]; + for (name, t, buffered) in cases { + let streamed = [(1, 0, t)]; + for join_type in [Inner, Left, Right, Full] { + for batch_size in [64, 333, 8192] { + for spill in [false, true] { + check_interval_join( + join_type, + &streamed, + buffered, + 1, + 128, + batch_size, + spill, + &format!( + "{name} {join_type} batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + Ok(()) +} + +/// Many streamed rows over key groups of various sizes, some larger than +/// the batch size: freezes cut through one streamed row's pairs and hold +/// the pairs of several rows, rows pass for none, one or several buffered +/// rows, and some keys exist on one side only. +#[tokio::test] +async fn join_filter_streamed_rows_across_freezes() -> Result<()> { + let mut buffered = vec![]; + for (key, rows, width) in [(0, 3, 1), (1, 150, 1), (2, 400, 2), (4, 1, 1), (5, 90, 4)] + { + let first_bid = buffered.len() as i32; + buffered.extend(interval_group(key, first_bid, rows, width)); + } + let mut state = 7u64; + let mut next = move || { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) as i64 + }; + let mut streamed = vec![]; + for (key, rows) in [(0, 2), (1, 7), (2, 5), (3, 2), (5, 9), (6, 1)] { + for _ in 0..rows { + let t = match next() % 6 { + 0 => None, + 1 => Some(-1 - next() % 10), + 2 => Some(5000 + next() % 10), + _ => Some(next() % 4100), + }; + streamed.push((key, streamed.len() as i32, t)); + } + } + for join_type in [Inner, Left, Right, Full] { + for (streamed_batch_rows, buffered_batch_rows) in [(4, 64), (1000, 37)] { + for batch_size in [50, 128, 1000] { + for spill in [false, true] { + check_interval_join( + join_type, + &streamed, + &buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + &format!( + "{join_type} streamed_batch_rows={streamed_batch_rows} \ + buffered_batch_rows={buffered_batch_rows} \ + batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + Ok(()) +} + +/// Each form of the interval filter lifts its subexpressions over one input +/// alone, and the bare columns it still reads stay input columns; the plain +/// form has nothing to lift. +#[test] +fn hoisted_join_filter_lifts_single_side_subexpressions() -> Result<()> { + use super::filter::{HoistedFilterInput, HoistedJoinFilter}; + use HoistedFilterInput::{Buffered, Column as Input, Streamed}; + let streamed = build_interval_streamed(&[(1, 0, Some(1))], 1).schema(); + let buffered = build_interval_buffered(&interval_group(1, 0, 1, 1), 1).schema(); + for streamed_side in [JoinSide::Left, JoinSide::Right] { + let buffered_side = streamed_side.negate(); + let hoist = |form| { + let filter = build_interval_filter(&streamed, &buffered, streamed_side, form); + HoistedJoinFilter::try_new(&filter, buffered_side, &streamed, &buffered) + }; + let to_strings = |exprs: &[PhysicalExprRef]| -> Vec { + exprs.iter().map(|e| e.to_string()).collect() + }; + assert!(hoist(IntervalFilter::Plain)?.is_none()); + + let hoisted = hoist(IntervalFilter::Buffered)?.unwrap(); + assert_eq!( + to_strings(&hoisted.buffered_exprs), + vec!["eff@2 + 0", "next_eff@3 + 0", "next_eff@3 - eff@2 > 0"] + ); + assert!(hoisted.streamed_exprs.is_empty()); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + Input(ColumnIndex { index: 2, side }), + Buffered(0), + Buffered(1), + Buffered(2), + ] if *side == streamed_side + )); + assert_eq!( + hoisted.expression.to_string(), + "t@0 > __buffered_0@1 AND t@0 <= __buffered_1@2 AND __buffered_2@3" + ); + + let hoisted = hoist(IntervalFilter::Streamed)?.unwrap(); + assert!(hoisted.buffered_exprs.is_empty()); + assert_eq!( + to_strings(&hoisted.streamed_exprs), + vec!["t@0 + 0", "t@0 - 0", "t@0 + 0 >= t@0 - 0"] + ); + assert_eq!(hoisted.streamed_projection, vec![2]); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + Input(ColumnIndex { index: 2, side: eff_side }), + Input(ColumnIndex { index: 3, side: next_eff_side }), + Streamed(0), + Streamed(1), + Streamed(2), + ] if *eff_side == buffered_side && *next_eff_side == buffered_side + )); + assert_eq!( + hoisted.expression.to_string(), + "__streamed_0@2 > eff@0 AND __streamed_1@3 <= next_eff@1 AND __streamed_2@4" + ); + + let hoisted = hoist(IntervalFilter::Both)?.unwrap(); + assert_eq!( + to_strings(&hoisted.buffered_exprs), + vec!["eff@2 + 0", "next_eff@3 + 0", "next_eff@3 - eff@2 > 0"] + ); + assert_eq!( + to_strings(&hoisted.streamed_exprs), + vec!["t@0 + 0", "t@0 - 0", "t@0 + 0 >= t@0 - 0"] + ); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + Streamed(0), + Buffered(0), + Streamed(1), + Buffered(1), + Buffered(2), + Streamed(2), + ] + )); + assert_eq!( + hoisted.expression.to_string(), + "__streamed_0@0 > __buffered_0@1 AND __streamed_1@2 <= __buffered_1@3 \ + AND __buffered_2@4 AND __streamed_2@5" + ); + } + Ok(()) +} + +/// Key groups larger than the batch size on both sides, with streamed rows +/// passing the filter for none, one or several buffered rows, and buffered +/// key groups no streamed row matches between them. A FULL join stages the +/// null-joined rows of such a group at the next freeze, which can fall +/// between two runs of one streamed row's pairs; the row must still be +/// emitted null-joined at most once. +#[tokio::test] +async fn join_filter_large_groups_on_both_sides() -> Result<()> { + let mut state = 11u64; + let mut next = move || { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) as i64 + }; + let mut buffered = vec![]; + let mut streamed = vec![]; + for (key, streamed_rows, buffered_rows, width) in [ + (0, 140, 140, 1), + (1, 0, 50, 1), + (2, 200, 70, 2), + (3, 0, 3, 1), + (4, 3, 140, 1), + (5, 140, 2, 1), + (6, 5, 0, 1), + (7, 70, 70, 1), + ] { + let first_bid = buffered.len() as i32; + buffered.extend(interval_group(key, first_bid, buffered_rows, width)); + for _ in 0..streamed_rows { + let t = match next() % 4 { + 0 => None, + 1 => Some(-1 - next() % 10), + _ => Some(next() % (buffered_rows as i64 * 12 + 1)), + }; + streamed.push((key, streamed.len() as i32, t)); + } + } + let never: Vec = streamed + .iter() + .map(|&(key, sid, _)| (key, sid, Some(-1))) + .collect(); + for (name, streamed) in [("mixed", &streamed), ("never", &never)] { + for join_type in [Inner, Left, Right, Full] { + for (streamed_batch_rows, buffered_batch_rows) in + [(64, 64), (1000, 37), (7, 1000)] + { + for batch_size in [64, 100, 8192] { + for spill in [false, true] { + check_interval_join( + join_type, + streamed, + &buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + &format!( + "{name} {join_type} streamed_batch_rows={streamed_batch_rows} \ + buffered_batch_rows={buffered_batch_rows} \ + batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + } + Ok(()) +} + /// Returns the column names on the schema fn columns(schema: &Schema) -> Vec { schema.fields().iter().map(|f| f.name().clone()).collect() diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 9c10c3e88ee..627be899e40 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -689,8 +689,8 @@ object CometConf extends ShimCometConf { "(or `k0,k1`, keeping c0), or spark, with `c0,k` for c0 + k*L ns per row, where L is " + "the number of leaf columns the class prices (for the functions of an aggregate or a " + "window, the number of functions). The classes are shuffleWrite, shuffleRead, sort, " + - "sortSpill, smj, smjCondition, bhj, predicate, projectPassThrough, expression, agg, " + - "aggObjectHash, aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, " + + "sortSpill, smj, smjCrossCondition, bhj, predicate, projectPassThrough, expression, " + + "agg, aggObjectHash, aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, " + "aggPercentileApprox, aggOther, aggArrayKey, codegenDispatch, window, " + "windowAggregate, windowOffset, windowRank, wglPartial, wglFinal, expand, generate, " + "rowLocal, the comet-only c2r and r2c, and the spark-only " + diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 0b66d06b94d..cef4dd8b2a7 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -25,7 +25,7 @@ import scala.collection.mutable import org.apache.spark.internal.Logging import org.apache.spark.sql.SparkSession -import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, Expression, NamedExpression, ScalaUDF} +import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, CaseWhen, Expression, If, NamedExpression, ScalaUDF} import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Partial} import org.apache.spark.sql.catalyst.plans.logical.statsEstimation.EstimationUtils import org.apache.spark.sql.catalyst.rules.Rule @@ -124,18 +124,14 @@ class EngineCostModel( private def aggregateShare(agg: BaseAggregateExec): Double = if (agg.aggregateExpressions.exists(_.mode == Complete)) 1.0 else 0.5 - /** Whether `condition` of `join`, run as `plan`, is one validity interval. */ - def validityInterval( - condition: Expression, - join: SortMergeJoinExec, - plan: SparkPlan): Boolean = { - val sides = if (plan.children.size == 2) plan.children else join.children - JoinConditionShape.isValidityInterval( - condition, - sides.head, - sides(1), - JoinConditionShape.aliases(plan.children ++ join.children)) - } + /** Whether `condition` of `join` holds a CASE or an IF over columns of both its inputs. */ + def crossInputConditional(condition: Expression, join: SortMergeJoinExec): Boolean = + condition.exists { + case e @ (_: CaseWhen | _: If) => + e.references.exists(join.left.outputSet.contains) && + e.references.exists(join.right.outputSet.contains) + case _ => false + } private def classTerms(costClass: CostClass, plan: SparkPlan, engine: Engine): Seq[Term] = { val op = sparkOperator(plan) @@ -184,9 +180,9 @@ class EngineCostModel( functions.distinct.map { c => Term(c, Width(functions.size, 0), share * functions.count(_ == c)) } - case (SmjCondition, join: SortMergeJoinExec) => - if (join.condition.exists(c => !validityInterval(c, join, plan))) { - Seq(Term(SmjCondition, out)) + case (SmjCrossCondition, join: SortMergeJoinExec) => + if (join.condition.exists(crossInputConditional(_, join))) { + Seq(Term(SmjCrossCondition, out)) } else { Nil } diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index d900df08d2d..3d2e4284526 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -172,7 +172,7 @@ object EngineCostTable { case object Sort extends CostClass("sort") case object SortSpill extends CostClass("sortSpill") case object Smj extends CostClass("smj") - case object SmjCondition extends CostClass("smjCondition") + case object SmjCrossCondition extends CostClass("smjCrossCondition") case object Bhj extends CostClass("bhj") case object Predicate extends CostClass("predicate") case object ProjectPassThrough extends CostClass("projectPassThrough") @@ -208,7 +208,7 @@ object EngineCostTable { Sort, SortSpill, Smj, - SmjCondition, + SmjCrossCondition, Bhj, Predicate, ProjectPassThrough, @@ -322,19 +322,20 @@ object EngineCostTable { * cannot be predicted. * - `smj`: the join over its sorted inputs, every output leaf, noisy. `bhj`: the probe side, * every output leaf; the nested `bhj` is noisy (Comet 0 to 350, Spark 50 to 4200 ns) and - * takes the flat line. `smjCondition`: what a join condition adds to `smj`, every output - * leaf. Comet joins every pair of rows of equal keys into a batch of all its output columns - * before the condition drops most of them; Spark tests the condition on the pair first. In - * a local micro-benchmark of left joins on a 14-day band (1 to 4% of 20 to 100 million - * pairs passing), Comet cost what the join without the condition cost, 0.44 to 1.2 us per - * output row and 21 ns per output leaf, Spark 3.6 to 4 times less (0.12 to 0.30 us, 4.5 ns - * per leaf); inner joins and small key groups ran 1 to 1.8 times slower. fbj_order_type's - * left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark. The - * pairs per output row are not estimated, so the line keeps the ratio of 3.5 at a price - * between the two, high enough to outweigh `smj` and the conversions around the join at any - * width. A condition that is one validity interval ([[JoinConditionShape]]) with both - * bounds from one row of one input, where Comet costs about what Spark does, adds nothing; - * bounds from two inputs joined below make a band again. + * takes the flat line. A join condition adds nothing to `smj`: the native join tests it on + * each pair of rows of equal keys before building the output, and on the cluster a left + * join under a validity interval or a band cost 5.7 to 22 ns per pair natively against 9.6 + * to 15 in Spark at 15 to 519 output leaves, an inner one 7.4 against 19.3, and one with + * 100 pairs per key about what Spark costs per output row. The exception is a condition + * holding a CASE or an IF over columns of both inputs, which adds `smjCrossCondition`, + * every output leaf: on the cluster, left joins with about one passing pair per streamed + * row under `CASE WHEN l.c = r.c THEN l.t - r.eff ELSE 10 - (r.nxt - l.t) END BETWEEN 1 AND + * 10` cost 28 to 51 ns per pair natively against 8.7 to 17.6 in Spark (1.7 to 4.5 times), + * at 10k and 100 keys and 32 and 128 pairs per key. Cross-input `datediff` cost about what + * Spark does, `instr` and an OR over a cast of one input 4 to 10 times less, and conditions + * over one input at most Spark's price. The pairs per output row are not estimated, so the + * line keeps the former `smjCondition`'s, Comet 3.5 times Spark at a price high enough to + * outweigh `smj` and the conversions around the join at any width. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when @@ -418,7 +419,7 @@ object EngineCostTable { (R2C, Flat) -> Line(3.4, 7.67, 0.036, 0, 0), (R2C, Nested) -> Line(0, 9.74, 0.040, 0, 0)) ++ both(ExprOverScan, Line(0, 0, 0, 2.5, 0)) ++ - both(SmjCondition, Line(3000, 80, 0, 850, 23)) ++ + both(SmjCrossCondition, Line(3000, 80, 0, 850, 23)) ++ both(AggObjectHash, Line(1000, 0, 0, 2500, 0)) ++ both(AggDeclarative, Line(6, 0, 0, 15, 0)) ++ both(AggDeclarativeNoCodegen, Line(0, 0, 0, 170, 0)) ++ @@ -481,7 +482,7 @@ object EngineCostTable { */ val operatorClasses: Map[String, Seq[CostClass]] = Map( "SortExec" -> Seq(Sort), - "SortMergeJoinExec" -> Seq(Smj, SmjCondition), + "SortMergeJoinExec" -> Seq(Smj, SmjCrossCondition), "BroadcastHashJoinExec" -> Seq(Bhj), "WindowExec" -> Seq(Window), "WindowGroupLimitExec" -> Seq(WglPartial, WglFinal), diff --git a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala deleted file mode 100644 index de0d07eff9d..00000000000 --- a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala +++ /dev/null @@ -1,218 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -package org.apache.comet.rules - -import scala.collection.mutable - -import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimestampAddYMInterval, TruncDate, TruncTimestamp, WindowExpression} -import org.apache.spark.sql.comet.CometExec -import org.apache.spark.sql.execution.{ProjectExec, SparkPlan, UnionExec} -import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, QueryStageExec} -import org.apache.spark.sql.execution.exchange.ReusedExchangeExec - -/** - * Recognizes a join condition that is a single validity interval: `L <= V < U` with `V` from one - * side, `L` and `U` from the other, `U` not derived from `L` by a constant offset, and `L` and - * `U` read from one row of one input of the other side, with no join or union between them. - */ -object JoinConditionShape { - - private case class Bound(small: Expression, big: Expression, nullCheck: Option[Expression]) - - private def unwrap(e: Expression): Expression = e match { - case c: Cast => unwrap(c.child) - case t: TruncTimestamp => unwrap(t.timestamp) - case t: TruncDate => unwrap(t.date) - case c: Coalesce if c.children.size > 1 && c.children.tail.forall(_.foldable) => - unwrap(c.children.head) - case other => other - } - - private val timeAddClasses = Set("TimeAdd", "TimestampAddInterval") - - private def offsetBase(e: Expression): Option[Expression] = e match { - case a: Add if a.right.foldable => Some(a.left) - case a: Add if a.left.foldable => Some(a.right) - case s: Subtract if s.right.foldable => Some(s.left) - case d: DateAdd if d.days.foldable => Some(d.startDate) - case d: DateSub if d.days.foldable => Some(d.startDate) - case t - if timeAddClasses(t.getClass.getSimpleName) && t.children.size >= 2 && - t.children(1).foldable => - Some(t.children.head) - case d: DateAddInterval if d.interval.foldable => Some(d.start) - case d: DateAddYMInterval if d.interval.foldable => Some(d.date) - case t: TimestampAddYMInterval if t.interval.foldable => Some(t.timestamp) - case m: AddMonths if m.numMonths.foldable => Some(m.startDate) - case _ => None - } - - private def strip(e: Expression): Expression = { - val u = unwrap(e) - offsetBase(u).map(strip).getOrElse(u) - } - - private def comparison(e: Expression): Option[(Expression, Expression)] = e match { - case LessThan(a, b) => Some((a, b)) - case LessThanOrEqual(a, b) => Some((a, b)) - case GreaterThan(a, b) => Some((b, a)) - case GreaterThanOrEqual(a, b) => Some((b, a)) - case _ => None - } - - private def bound(e: Expression): Option[Bound] = e match { - case Or(IsNull(x), c) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } - case Or(c, IsNull(x)) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } - case c => comparison(c).map { case (s, b) => Bound(s, b, None) } - } - - private def conjuncts(e: Expression): Seq[Expression] = e match { - case And(a, b) => conjuncts(a) ++ conjuncts(b) - case other => Seq(other) - } - - private def same(a: Expression, b: Expression): Boolean = unwrap(a).semanticEquals(unwrap(b)) - - /** The aliases defined in `plans` and below, by the id of the attribute each defines. */ - def aliases(plans: Seq[SparkPlan]): Map[ExprId, Expression] = { - val result = mutable.Map[ExprId, Expression]() - def visit(node: SparkPlan): Unit = { - val operators = node match { - case c: CometExec => Seq(c, c.originalPlan) - case other => Seq(other) - } - operators - .flatMap(_.expressions) - .foreach(_.foreach { - case a: Alias => result.getOrElseUpdate(a.exprId, a.child) - case _ => - }) - node match { - case s: QueryStageExec => visit(s.plan) - case _ => node.children.foreach(visit) - } - } - plans.foreach(visit) - result.toMap - } - - private def sources(e: Expression, aliases: Map[ExprId, Expression], depth: Int): Set[ExprId] = - strip(e) match { - case a: Attribute => - aliases.get(a.exprId) match { - case Some(_: WindowExpression) => Set(a.exprId) - case Some(child) if depth < 64 => sources(child, aliases, depth + 1) - case _ => Set(a.exprId) - } - case other => other.references.map(_.exprId).toSet - } - - private def inputs(node: SparkPlan): Seq[SparkPlan] = node match { - case s: QueryStageExec => Seq(s.plan) - case a: AdaptiveSparkPlanExec => Seq(a.executedPlan) - case other => other.children - } - - private def original(node: SparkPlan): SparkPlan = node match { - case c: CometExec => c.originalPlan - case other => other - } - - private def outputs(node: SparkPlan, id: ExprId): Boolean = node.output.exists(_.exprId == id) - - private def producers(node: SparkPlan, id: ExprId, depth: Int): Option[Seq[List[SparkPlan]]] = - if (depth > 256 || !outputs(node, id)) { - None - } else { - node match { - case r: ReusedExchangeExec => - val i = r.output.indexWhere(_.exprId == id) - producers(r.child, r.child.output(i).exprId, depth + 1).map(_.map(r :: _)) - case _ => - inputs(node).find(outputs(_, id)) match { - case Some(child) => producers(child, id, depth + 1).map(_.map(node :: _)) - case None => - val projected = original(node) match { - case p: ProjectExec => - p.projectList.collectFirst { case a: Alias if a.exprId == id => a.child } - case _ => None - } - (projected, inputs(node)) match { - case (Some(e), Seq(child)) if e.references.nonEmpty => - val traced = e.references.toSeq.map(a => producers(child, a.exprId, depth + 1)) - if (traced.forall(_.isDefined)) { - Some(traced.flatMap(_.get).map(node :: _)) - } else { - None - } - case _ => Some(Seq(List(node))) - } - } - } - } - - private def oneRow(side: SparkPlan, refs: AttributeSet): Boolean = { - val traced = refs.toSeq.map(a => producers(side, a.exprId, 0)) - traced.nonEmpty && traced.forall(_.isDefined) && { - val paths = traced.flatMap(_.get) - val common = (0 until paths.map(_.size).min) - .takeWhile(i => paths.forall(_(i) eq paths.head(i))) - .size - val below = paths.map(_.drop(common)) - val heads = below.flatMap(_.headOption) - heads.forall(_ eq heads.head) && - paths.forall(_.forall(n => !original(n).isInstanceOf[UnionExec])) && - below.forall(_.forall(n => inputs(n).size <= 1)) - } - } - - /** - * Whether `condition` of a join between `left` and `right` is one validity interval whose - * bounds `aliases` does not derive from one another and that read one row of one input. - */ - def isValidityInterval( - condition: Expression, - left: SparkPlan, - right: SparkPlan, - aliases: Map[ExprId, Expression]): Boolean = { - conjuncts(condition).map(bound) match { - case Seq(Some(first), Some(second)) => - val shapes = Seq((first, second), (second, first)).collect { - case (lower, upper) if same(lower.big, upper.small) && lower.nullCheck.isEmpty => - (upper.small, lower.small, upper.big, upper.nullCheck) - } - shapes.exists { case (v, l, u, nullCheck) => - val (vRefs, lRefs, uRefs) = (v.references, l.references, u.references) - def onOtherSides(side: SparkPlan, other: SparkPlan): Boolean = - vRefs.subsetOf(side.outputSet) && lRefs.subsetOf(other.outputSet) && - uRefs.subsetOf(other.outputSet) - val boundsSide = - if (onOtherSides(left, right)) Some(right) - else if (onOtherSides(right, left)) Some(left) - else None - vRefs.nonEmpty && lRefs.nonEmpty && uRefs.nonEmpty && - nullCheck.forall(same(_, u)) && - sources(l, aliases, 0).intersect(sources(u, aliases, 0)).isEmpty && - boundsSide.exists(oneRow(_, lRefs ++ uRefs)) - } - case _ => false - } - } -} diff --git a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala index f18689c51da..8bb8d5ddf70 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala @@ -1142,6 +1142,62 @@ class CometJoinSuite extends CometTestBase { } } + test("SortMergeJoin with a validity-interval join filter over wide rows and large key groups") { + withTempPath { dir => + withSQLConf( + CometConf.COMET_EXEC_SORT_MERGE_JOIN_ENABLED.key -> "true", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_WITH_JOIN_FILTER_ENABLED.key -> "true", + CometConf.COMET_BATCH_SIZE.key -> "1000", + SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + // Each currency but the last has 3000 consecutive validity intervals, so an + // order meets a key group three times the batch size and passes for one + // interval (several for the overlapping currency 3), or for none when its + // time is out of range or NULL. Currency 4 has no rates, currency 5 no orders. + val payload = (0 until 30).map { c => + if (c % 2 == 0) s"id * $c AS p$c" else s"concat('payload_$c', '_', id) AS p$c" + } + val ordersPath = s"${dir.getCanonicalPath}/orders" + spark + .range(0, 600, 1, 1) + .selectExpr(Seq( + "id", + "CAST(id % 5 AS INT) AS cur", + "CASE WHEN id % 11 = 0 THEN NULL WHEN id % 13 = 0 THEN -5 " + + "WHEN id % 17 = 0 THEN 50000 ELSE (id * 7919) % 30000 + 1 END AS t") ++ + payload: _*) + .write + .parquet(ordersPath) + val ratesPath = s"${dir.getCanonicalPath}/rates" + spark + .range(0, 15000, 1, 1) + .selectExpr( + "CAST(CASE WHEN id < 12000 THEN id DIV 3000 ELSE 5 END AS INT) AS cur", + "(id % 3000) * 10 AS eff", + "(id % 3000) * 10 + CASE WHEN id DIV 3000 = 3 THEN 30 ELSE 10 END AS next_eff", + "concat('rate_', id) AS rate") + .write + .parquet(ratesPath) + + withParquetTable(ordersPath, "orders") { + withParquetTable(ratesPath, "rates") { + val condition = "o.cur = r.cur AND o.t > r.eff AND o.t <= r.next_eff" + for (joinType <- Seq("INNER", "LEFT", "FULL")) { + checkSparkAnswerAndOperator( + sql(s"SELECT o.*, r.* FROM orders o $joinType JOIN rates r ON $condition")) + } + checkSparkAnswerAndOperator( + sql(s"SELECT o.*, r.* FROM rates r RIGHT JOIN orders o ON $condition")) + checkSparkAnswerAndOperator( + sql("SELECT o.id, count(r.rate), sum(o.p28) " + + s"FROM orders o LEFT JOIN rates r ON $condition GROUP BY o.id")) + } + } + } + } + } + test("full outer join") { withTempView("`left`", "`right`", "allNulls") { allNulls.createOrReplaceTempView("allNulls") diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala new file mode 100644 index 00000000000..496bfd3c521 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -0,0 +1,823 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.exec + +import java.io.File +import java.nio.file.Files +import java.sql.{Date, Timestamp} +import java.time.{Instant, LocalDate, LocalDateTime, ZoneOffset} +import java.time.format.DateTimeFormatter + +import scala.collection.mutable +import scala.util.Random + +import org.apache.commons.io.FileUtils +import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.comet.CometSortMergeJoinExec +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.joins.SortMergeJoinExec +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types._ + +import org.apache.comet.CometConf + +/** + * Differential tests of Comet's sort-merge join with a join filter against Spark, over generated + * data with key groups of one row, about a batch and several batches on either side, NULL keys + * and NULL filter columns, and a wide streamed side with strings, decimals and nested columns. + * + * Every join type runs with filters of different shapes and selectivities (a validity interval + * bounded by one side, a band bounded by the other, one-sided conditions, always true and always + * false, almost nothing or almost everything passing, casts of one or both sides) at several + * batch sizes and under a memory pool small enough to make the join spill. Results are compared + * with Spark as multisets. + * + * The default run covers each join type and filter at one batch size and a subset at the others. + * The full matrix runs only with `-Dcomet.test.smjFuzz.full=true`; + * `-Dcomet.test.smjFuzz.seed=` changes the data seed. + */ +class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHelper { + + private val seed: Long = + sys.props.get("comet.test.smjFuzz.seed").map(_.toLong).getOrElse(20261007L) + private val fullMatrix: Boolean = sys.props.get("comet.test.smjFuzz.full").contains("true") + + // ---------------------------------------------------------------- data generation + + private val leftSchema = StructType( + Seq( + StructField("id", LongType, nullable = false), + StructField("k", IntegerType), + StructField("kc", StringType), + StructField("t", LongType), + StructField("w", LongType), + StructField("lv", IntegerType), + StructField("s", StringType), + StructField("d38", DecimalType(38, 10)), + StructField("d10", DecimalType(10, 2)), + StructField("dbl", DoubleType), + StructField("ts", TimestampType), + StructField("dt", DateType), + StructField("b", BooleanType), + StructField( + "st", + StructType( + Seq( + StructField("a", IntegerType), + StructField("b", StringType), + StructField("c", ArrayType(IntegerType))))), + StructField("arr", ArrayType(StringType)), + StructField("m", MapType(StringType, IntegerType)), + StructField("ss", StringType), + StructField("li", StringType), + StructField("k2", IntegerType), + StructField("flag", BooleanType), + StructField("ldt", DateType))) + + private val rightSchema = StructType( + Seq( + StructField("id", LongType, nullable = false), + StructField("k", IntegerType), + StructField("kc", StringType), + StructField("eff", LongType), + StructField("next_eff", LongType), + StructField("rv", IntegerType), + StructField("rs", StringType), + StructField("rdec", DecimalType(18, 4)), + StructField("rarr", ArrayType(IntegerType)), + StructField( + "rst", + StructType(Seq(StructField("x", DoubleType), StructField("y", StringType)))), + StructField("rt", TimestampType), + StructField("rss", StringType), + StructField("rk2", IntegerType), + StructField("rdt", DateType))) + + private case class DataSet(name: String, left: String, right: String, stats: String) + + private var tempRoot: File = _ + private val dataSets = mutable.Map[String, DataSet]() + + override protected def afterAll(): Unit = { + try { + if (tempRoot != null) FileUtils.deleteQuietly(tempRoot) + } finally { + super.afterAll() + } + } + + private def withConfs[T](pairs: (String, String)*)(f: => T): T = { + var result: Option[T] = None + withSQLConf(pairs: _*) { result = Some(f) } + result.get + } + + private def pick[T](r: Random, xs: Seq[T]): T = xs(r.nextInt(xs.length)) + + private def orNull(r: Random, pct: Int)(v: => Any): Any = if (r.nextInt(100) < pct) null else v + + private def randomDecimal(r: Random, precision: Int, scale: Int): java.math.BigDecimal = { + val digits = (0 until 1 + r.nextInt(precision)).map(_ => ('0' + r.nextInt(10)).toChar) + val unscaled = new java.math.BigInteger(digits.mkString) + new java.math.BigDecimal(if (r.nextBoolean()) unscaled.negate() else unscaled, scale) + } + + private def randomString(r: Random): String = r.nextInt(30) match { + case 0 => "" + case 1 => "ж€😀 ünï" + case 2 => r.alphanumeric.take(1500 + r.nextInt(1000)).mkString + case _ => r.alphanumeric.take(1 + r.nextInt(12)).mkString + } + + private def randomDouble(r: Random): Double = + if (r.nextInt(4) == 0) { + pick(r, Seq(Double.NaN, -0.0d, 0.0d, Double.PositiveInfinity, Double.MinPositiveValue)) + } else r.nextDouble() * 2e6 - 1e6 + + private case class KeyGroup(k: Int, nl: Int, nr: Int, base: Long, mode: Int) + + private def leftRow(r: Random, g: KeyGroup): Seq[Any] = { + val k: Any = if (g.k < 0) null else g.k + val kc = + if (g.k < 0) "c0" else if (g.k % 23 == 0 && r.nextInt(3) == 0) null else s"c${g.k % 5}" + val span = math.max(g.nr, 1) * 10L + 40 + Seq( + k, + kc, + orNull(r, 8)(g.base - 20 + (r.nextLong() & Long.MaxValue) % span), + orNull(r, 5)(pick(r, Seq(0L, 5L, 15L, 30L))), + orNull(r, 10)(r.nextInt(100)), + orNull(r, 10)(randomString(r)), + orNull(r, 10)(randomDecimal(r, 38, 10)), + orNull(r, 10)(randomDecimal(r, 10, 2)), + orNull(r, 10)(randomDouble(r)), + orNull(r, 10)( + Timestamp.from(Instant.ofEpochSecond(r.nextInt(2000000000) * 2L - 1500000000L, 1000L))), + orNull(r, 10)(Date.valueOf(LocalDate.ofEpochDay(r.nextInt(60000) - 20000L))), + orNull(r, 10)(r.nextBoolean()), + orNull(r, 10)( + Row( + orNull(r, 20)(r.nextInt(1000)), + orNull(r, 20)(randomString(r)), + orNull(r, 20)(Seq.fill(r.nextInt(4))(orNull(r, 20)(r.nextInt(50)))))), + orNull(r, 10)( + Seq.fill(r.nextInt(5))(orNull(r, 20)(r.alphanumeric.take(r.nextInt(8)).mkString))), + orNull(r, 10)((0 until r.nextInt(3)).map(i => s"k$i" -> orNull(r, 20)(r.nextInt(9))).toMap)) + } + + private def rightRow(r: Random, g: KeyGroup, j: Int): Seq[Any] = { + val k: Any = if (g.k < 0) null else g.k + val kc = if (g.k < 0) "c0" else s"c${g.k % 5}" + val eff = g.base + j * 10L + val width = g.mode match { + case 0 => 10L + case 1 => 25L + case _ => 5L + } + Seq( + k, + kc, + orNull(r, 4)(eff), + orNull(r, 4)(eff + width), + orNull(r, 10)(r.nextInt(100)), + orNull(r, 10)(randomString(r)), + orNull(r, 10)(randomDecimal(r, 18, 4)), + orNull(r, 10)(Seq.fill(r.nextInt(4))(orNull(r, 20)(r.nextInt(50)))), + orNull(r, 10)(Row(orNull(r, 20)(randomDouble(r)), orNull(r, 20)(randomString(r))))) + } + + private val tsFormat = DateTimeFormatter.ofPattern("yyyy-MM-dd HH:mm:ss") + + private def formatSeconds(sec: Long): String = + LocalDateTime.ofEpochSecond(sec, 0, ZoneOffset.UTC).format(tsFormat) + + // A string column cast to TIMESTAMP in the filters: mostly valid timestamps near the key + // group's values, but also NULL, empty, malformed, out of range, date-only, padded and + // year-only strings, which Legacy casts turn into NULL or partial values. + private def timestampString(r: Random, sec: Long): Any = r.nextInt(24) match { + case 0 | 1 | 2 => null + case 3 => "" + case 4 => "not a timestamp" + case 5 => "2026-13-45 99:00:00" + case 6 => formatSeconds(sec).replace(' ', 'T').dropRight(3) + ":75" + case 7 => " " + formatSeconds(sec) + " " + case 8 => formatSeconds(sec).take(10) + case 9 => formatSeconds(sec).take(4) + case 10 => String.valueOf(sec) + case 11 => "ж€😀" + case _ => formatSeconds(sec) + } + + // A string column cast to INT in the filters: numbers, padded numbers, fractions, NULL and + // strings that Legacy casts turn into NULL. + private def intString(r: Random): Any = r.nextInt(16) match { + case 0 | 1 => null + case 2 => "" + case 3 => "x1" + case 4 => " " + r.nextInt(100) + " " + case 5 => s"${r.nextInt(100)}.${r.nextInt(10)}" + case 6 => "99999999999" + case _ => String.valueOf(r.nextInt(100)) + } + + private def leftExtra(r: Random, g: KeyGroup): Seq[Any] = { + val span = math.max(g.nr, 1) * 10L + 40 + val sec = g.base - 20 + r.nextInt(span.toInt) + Seq( + timestampString(r, sec), + intString(r), + orNull(r, 10)(r.nextInt(6)), + orNull(r, 15)(r.nextBoolean()), + orNull(r, 30)(Date.valueOf(LocalDate.ofEpochDay(sec / 86400 + r.nextInt(3) - 1)))) + } + + private def rightExtra(r: Random, g: KeyGroup, j: Int): Seq[Any] = { + val sec = g.base + j * 10L + r.nextInt(31) - 15 + Seq( + orNull(r, 10)(Timestamp.from(Instant.ofEpochSecond(sec))), + timestampString(r, sec + r.nextInt(21) - 10), + orNull(r, 10)(r.nextInt(6)), + orNull(r, 30)(Date.valueOf(LocalDate.ofEpochDay(sec / 86400 + r.nextInt(3) - 1)))) + } + + private val mainCombos: Seq[(Int, Int)] = Seq( + 1 -> 1, + 1 -> 7, + 7 -> 1, + 7 -> 7, + 8 -> 6, + 6 -> 8, + 63 -> 64, + 64 -> 64, + 65 -> 1, + 1 -> 65, + 130 -> 2, + 2 -> 130, + 140 -> 140, + 200 -> 70, + 1023 -> 1, + 1 -> 1024, + 1025 -> 3, + 3 -> 1025, + 2100 -> 1, + 1 -> 2100, + 0 -> 10, + 10 -> 0, + 0 -> 1100, + 1100 -> 0) + + private val bigCombos: Seq[(Int, Int)] = + Seq(5000 -> 2, 2 -> 5000, 3000 -> 40, 40 -> 3000, 600 -> 600, 9000 -> 1, 1 -> 9000) + + private def dataSet(name: String): DataSet = synchronized { + dataSets.getOrElseUpdate( + name, + name match { + case "main" => createDataSet(name, seed, mainCombos, 300) + case "alt" => createDataSet(name, seed + 1, mainCombos, 300) + case "big" => createDataSet(name, seed, bigCombos, 200) + }) + } + + private def createDataSet( + name: String, + dataSeed: Long, + combos: Seq[(Int, Int)], + numRandom: Int): DataSet = { + val r = new Random(dataSeed ^ name.hashCode) + def small(): Int = r.nextInt(10) match { + case 0 => 0 + case 1 | 2 | 3 => 1 + case 9 => 6 + r.nextInt(15) + case _ => 2 + r.nextInt(4) + } + val sized = r.shuffle(combos ++ Seq.fill(numRandom)((small(), small()))) + val groups = sized.zipWithIndex.map { case ((nl, nr), i) => + KeyGroup(i * 3 + r.nextInt(3), nl, nr, r.nextInt(100000) * 10L, r.nextInt(3)) + } ++ Seq(KeyGroup(-1, 40 + r.nextInt(40), 40 + r.nextInt(40), 5000L, 0)) + val lx = new Random(dataSeed ^ name.hashCode ^ 0x5eed1L) + val rx = new Random(dataSeed ^ name.hashCode ^ 0x5eed2L) + val leftRows = + r.shuffle(groups.flatMap(g => Seq.fill(g.nl)(leftRow(r, g) ++ leftExtra(lx, g)))) + val rightRows = r.shuffle( + groups.flatMap(g => (0 until g.nr).map(j => rightRow(r, g, j) ++ rightExtra(rx, g, j)))) + + if (tempRoot == null) tempRoot = Files.createTempDirectory("comet-smj-fuzz").toFile + val leftPath = new File(tempRoot, s"$name-left").getCanonicalPath + val rightPath = new File(tempRoot, s"$name-right").getCanonicalPath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val lr = leftRows.zipWithIndex.map { case (v, i) => Row.fromSeq(i.toLong +: v) } + val rr = rightRows.zipWithIndex.map { case (v, i) => Row.fromSeq(i.toLong +: v) } + spark + .createDataFrame(spark.sparkContext.parallelize(lr, 4), leftSchema) + .write + .option("parquet.block.size", 65536) + .parquet(leftPath) + spark + .createDataFrame(spark.sparkContext.parallelize(rr, 4), rightSchema) + .write + .option("parquet.block.size", 65536) + .parquet(rightPath) + } + val ds = DataSet( + name, + s"smj_fuzz_${name}_l", + s"smj_fuzz_${name}_r", + s"left rows ${leftRows.size}, right rows ${rightRows.size}, keys ${groups.size}, " + + s"max group ${groups.map(_.nl).max}/${groups.map(_.nr).max}, " + + s"key pairs ${groups.filter(_.k >= 0).map(g => g.nl.toLong * g.nr).sum}") + spark.read.parquet(leftPath).createOrReplaceTempView(ds.left) + spark.read.parquet(rightPath).createOrReplaceTempView(ds.right) + ds + } + + // ---------------------------------------------------------------- query shapes + + private case class JoinKind(name: String, sql: String, existence: Boolean = false) + + private val inner = JoinKind("inner", "INNER JOIN") + private val leftOuter = JoinKind("left", "LEFT JOIN") + private val rightOuter = JoinKind("right", "RIGHT JOIN") + private val fullOuter = JoinKind("full", "FULL JOIN") + private val leftSemi = JoinKind("semi", "LEFT SEMI JOIN", existence = true) + private val leftAnti = JoinKind("anti", "LEFT ANTI JOIN", existence = true) + private val joinKinds = Seq(inner, leftOuter, rightOuter, fullOuter, leftSemi, leftAnti) + + private case class Filter(name: String, sql: String) + + private val interval = Filter("interval", "l.t > r.eff AND l.t <= r.next_eff") + private val band = Filter("band", "r.eff >= l.t - l.w AND r.eff < l.t + l.w") + private val leftOnly = Filter("left_only", "l.lv % 3 = 0") + private val rightOnly = Filter("right_only", "r.rv % 3 = 0") + private val alwaysFalse = Filter("false", "l.id + r.id < 0") + private val alwaysTrue = Filter("true", "l.id + r.id >= 0") + private val rare = Filter("rare", "(l.id * 31 + r.id * 17) % 997 = 0") + private val most = Filter("most", "(l.id + r.id) % 10 <> 0") + private val nullable = Filter("nullable_cmp", "l.lv < r.rv") + private val typed = Filter("typed", "l.d10 * 2 >= r.rdec OR l.s < r.rs") + private val leftCast = + Filter("left_cast", "CAST(CAST(l.t AS STRING) AS BIGINT) > r.eff AND l.t <= r.next_eff") + private val leftTimestampCast = Filter( + "left_ts_cast", + "CAST(CAST(l.ts AS STRING) AS TIMESTAMP) < CAST(r.eff * 2000 AS TIMESTAMP) AND l.t > r.eff") + private val bothCast = Filter( + "both_cast", + "CAST(CAST(l.t AS STRING) AS BIGINT) > CAST(CAST(r.eff AS STRING) AS BIGINT) AND " + + "CAST(l.t AS DECIMAL(20, 0)) <= CAST(r.next_eff AS DECIMAL(20, 0))") + private val filters = + Seq( + interval, + band, + leftOnly, + rightOnly, + alwaysFalse, + alwaysTrue, + rare, + most, + nullable, + typed, + leftCast, + leftTimestampCast, + bothCast) + + // Boolean and conditional operators around subexpressions that read one side only, which + // the native join lifts out of the filter and evaluates once per row: OR, CASE, IF, NOT, + // IS [NOT] NULL, IN and NOT IN, COALESCE and null-safe equality, with casts of strings that + // are often invalid, so the lifted parts are often NULL and three-valued logic matters. + private val lts = "CAST(l.ss AS TIMESTAMP)" + private val rts = "CAST(r.rss AS TIMESTAMP)" + private val lint = "CAST(l.li AS INT)" + private val orTs = Filter("or_ts", s"l.t > r.eff OR $lts < r.rt") + private val orTsNull = Filter("or_ts_null", s"$lts < r.rt OR l.lv IS NULL") + private val orMostlyTrue = + Filter("or_mostly_true", s"(l.id + r.id) % 10 <> 0 OR $lts >= r.rt") + private val orBufferedLifted = + Filter("or_buffered_lifted", s"l.lv < r.rv OR $rts > CAST(l.t AS TIMESTAMP)") + private val orSingleSides = + Filter("or_single_sides", s"(l.flag AND $lts IS NOT NULL) OR r.rv % 3 = 0") + private val orBothLifted = + Filter("or_both_lifted", s"$lts < $rts OR ($lts IS NULL AND $rts IS NULL)") + private val caseFlag = + Filter("case_flag", s"CASE WHEN l.flag THEN $lts < r.rt ELSE r.eff <= l.t END") + private val caseRightWhen = Filter( + "case_right_when", + s"CASE WHEN r.rv % 2 = 0 THEN $lts < TIMESTAMP '1970-01-06 00:00:00' " + + s"WHEN r.rv IS NULL THEN l.flag ELSE $lint > 50 END") + private val caseLeftWhen = Filter( + "case_left_when", + s"CASE WHEN l.flag THEN $rts > TIMESTAMP '1970-01-06 00:00:00' ELSE r.rk2 > 2 END") + private val caseNested = Filter( + "case_nested", + "CASE WHEN l.k2 IS NULL THEN r.rv IS NULL WHEN l.k2 > 2 THEN " + + s"CASE WHEN $lts IS NULL THEN r.rk2 = l.k2 ELSE $lts < r.rt END " + + s"ELSE CASE WHEN r.rk2 IS NULL THEN l.flag ELSE $lint < r.rv END END") + private val caseValue = + Filter("case_value", s"CASE WHEN l.flag THEN $lts ELSE CAST(l.t AS TIMESTAMP) END < r.rt") + private val ifNull = + Filter("if_null", s"IF($lint IS NULL, r.rk2 > 2, $lint <= r.rv)") + private val notAndUpper = Filter("not_and_upper", "NOT (l.lv > r.rv AND upper(l.s) = r.rs)") + private val notOr = Filter("not_or", s"NOT ($lts >= r.rt OR l.flag)") + private val isNullMix = + Filter("is_null_mix", s"($lts IS NULL) = (r.rt IS NULL) OR ($lint + r.rv) IS NULL") + private val isNotNullCmp = + Filter("is_not_null_cmp", s"($lts < r.rt) IS NOT NULL AND NOT ($lts < r.rt)") + private val inList = + Filter("in_list", "l.k2 IN (1, 2, 3) AND l.t > r.eff OR r.rk2 IN (0, 4) AND l.lv < r.rv") + private val notInList = Filter( + "not_in_list", + s"l.k2 NOT IN (1, 3) AND r.rk2 NOT IN (2) OR $lint NOT IN (1, 2, 50) AND l.t <= r.next_eff") + private val inCross = + Filter("in_cross", s"r.rk2 IN (l.k2, l.k2 + 1, $lint) OR l.k2 IN (r.rk2, 3)") + private val coalesceTs = + Filter("coalesce_ts", s"coalesce($lts, r.rt) > CAST(r.eff AS TIMESTAMP)") + private val coalesceDate = + Filter("coalesce_date", "coalesce(l.ldt, r.rdt) >= CAST(r.rt AS DATE)") + private val threeValued = + Filter("three_valued", s"($lint > r.rv OR $lts < r.rt) AND NOT ($rts > $lts AND l.flag)") + private val notNullPart = + Filter("not_null_part", "NOT (CAST(l.ss AS INT) > 0) OR r.rv > 50") + private val nullSafeEq = + Filter("null_safe_eq", s"$lts <=> r.rt OR $lint <=> r.rk2") + private val liftedTwice = + Filter("lifted_twice", s"$lts >= r.rt AND $lts <= CAST(r.eff + 20 AS TIMESTAMP)") + private val upperInvalid = + Filter("upper_invalid", "upper(l.ss) = upper(r.rss) OR upper(l.s) < r.rs") + private val boolFilters = Seq( + orTs, + orTsNull, + orMostlyTrue, + orBufferedLifted, + orSingleSides, + orBothLifted, + caseFlag, + caseRightWhen, + caseLeftWhen, + caseNested, + caseValue, + ifNull, + notAndUpper, + notOr, + isNullMix, + isNotNullCmp, + inList, + notInList, + inCross, + coalesceTs, + coalesceDate, + threeValued, + notNullPart, + nullSafeEq, + liftedTwice, + upperInvalid) + private val defaultBoolFilters = + Seq(orTs, caseFlag, notAndUpper, inList, coalesceTs, threeValued) + + private case class Shape(name: String, confs: Seq[(String, String)], spill: Boolean = false) + + private def batch(n: Int) = CometConf.COMET_BATCH_SIZE.key -> n.toString + private def aqe(on: Boolean) = SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> on.toString + private def partitions(n: Int) = SQLConf.SHUFFLE_PARTITIONS.key -> n.toString + private val tinyPool = CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.0005" + private val manySplits = Seq( + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "65536", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1") + + private val b7 = Shape("batch7", Seq(batch(7), aqe(false), partitions(1))) + private val b64 = Shape("batch64", Seq(batch(64), aqe(true), partitions(2))) + private val b1024 = Shape("batch1024", Seq(batch(1024), aqe(true), partitions(2))) + private val b8192 = Shape("batch8192", Seq(batch(8192), aqe(false), partitions(3))) + private val spill64 = + Shape( + "spill64", + Seq(batch(64), aqe(true), partitions(1), tinyPool) ++ manySplits, + spill = true) + private val spill7 = + Shape( + "spill7", + Seq(batch(7), aqe(false), partitions(2), tinyPool) ++ manySplits, + spill = true) + private val allShapes = Seq(b7, b64, b1024, b8192, spill7, spill64) + + private val baseConfs: Seq[(String, String)] = Seq( + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.PREFER_SORTMERGEJOIN.key -> "true", + SQLConf.ANSI_ENABLED.key -> "false", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_ENABLED.key -> "true", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_WITH_JOIN_FILTER_ENABLED.key -> "true", + CometConf.COMET_FORCE_SHJ.key -> "false") + + private case class Q( + label: String, + sql: String, + expectNative: Boolean = true, + expectCondition: Boolean = true) + + private def joinQuery( + ds: DataSet, + kind: JoinKind, + filter: Filter, + flipped: Boolean, + twoKeys: Boolean): Q = { + val on = (if (twoKeys) "l.k = r.k AND l.kc = r.kc" else "l.k = r.k") + s" AND (${filter.sql})" + val (from, select) = + if (flipped) (s"${ds.right} r ${kind.sql} ${ds.left} l", if (kind.existence) "r.*" else "*") + else (s"${ds.left} l ${kind.sql} ${ds.right} r", if (kind.existence) "l.*" else "*") + val pushable = (filter, kind.name, flipped) match { + case (`leftOnly`, "inner" | "right", false) => true + case (`leftOnly`, "inner" | "left", true) => true + case (`leftOnly`, "semi" | "anti", true) => true + case (`leftOnly`, "semi", false) => true + case (`rightOnly`, "inner" | "left", false) => true + case (`rightOnly`, "inner" | "right", true) => true + case (`rightOnly`, "semi" | "anti", false) => true + case (`rightOnly`, "semi", true) => true + case _ => false + } + val label = s"${kind.name}${if (flipped) "_rl" else ""}${if (twoKeys) "_2keys" else ""}" + + s"/${filter.name}" + Q(label, s"SELECT $select FROM $from ON $on", expectCondition = !pushable) + } + + private def existenceQueries(ds: DataSet): Seq[Q] = Seq( + Q( + "exists_or", + s"SELECT l.* FROM ${ds.left} l WHERE l.lv > 90 OR EXISTS (SELECT 1 FROM ${ds.right} r " + + "WHERE r.k = l.k AND l.t > r.eff AND l.t <= r.next_eff)", + expectNative = false), + Q( + "in_or", + s"SELECT l.* FROM ${ds.left} l WHERE l.lv > 90 OR l.k IN (SELECT r.k FROM ${ds.right} r " + + "WHERE l.t > r.eff AND l.t <= r.next_eff)", + expectNative = false), + Q( + "not_exists", + s"SELECT l.* FROM ${ds.left} l WHERE NOT EXISTS (SELECT 1 FROM ${ds.right} r " + + "WHERE r.k = l.k AND r.eff >= l.t - l.w AND r.eff < l.t + l.w)")) + + // ---------------------------------------------------------------- comparison + + private def norm(v: Any): Any = v match { + case null => null + case b: Array[Byte] => ("bin", b.toList) + case d: Double => ("d", java.lang.Double.doubleToLongBits(d)) + case f: Float => ("f", java.lang.Float.floatToIntBits(f)) + case r: Row => r.toSeq.map(norm).toList + case m: scala.collection.Map[_, _] => m.map { case (k, x) => (norm(k), norm(x)) }.toMap + case s: scala.collection.Seq[_] => s.map(norm).toList + case other => other + } + + private def counts(rows: Seq[Row]): Map[Any, Int] = + rows.groupBy(r => norm(r)).map { case (k, v) => k -> v.size } + + private def showRow(v: Any): String = { + val s = String.valueOf(v) + if (s.length > 300) s.take(300) + "..." else s + } + + private case class Outcome( + label: String, + shape: String, + data: String, + sparkRows: Int, + cometRows: Int, + native: Boolean, + condition: Boolean, + spills: Long, + error: Option[String]) { + def line: String = + f"SMJFUZZ ${if (error.isEmpty) "PASS" else "FAIL"} seed=$seed data=$data%-4s " + + f"shape=$shape%-9s case=$label%-28s spark=$sparkRows%7d comet=$cometRows%7d " + + s"native=$native cond=$condition spills=$spills" + } + + private def cometJoins(plan: SparkPlan): Seq[CometSortMergeJoinExec] = + collect(plan) { case j: CometSortMergeJoinExec => j } + + private def runCase(ds: DataSet, shape: Shape, q: Q): Outcome = { + val ctx = s"[seed=$seed data=${ds.name} shape=${shape.name} case=${q.label}]" + withConfs(baseConfs ++ shape.confs: _*) { + val expected = withConfs(CometConf.COMET_ENABLED.key -> "false") { + sql(q.sql).collect().toSeq + } + val df = sql(q.sql) + val actual = + try Right(df.collect().toSeq) + catch { case e: Throwable => Left(e) } + val plan = df.queryExecution.executedPlan + val joins = cometJoins(plan) + val sparkJoins = collect(plan) { case j: SortMergeJoinExec => j } + val native = joins.nonEmpty && sparkJoins.isEmpty + val condition = joins.exists(_.condition.isDefined) + val spills = joins.flatMap(_.metrics.get("spill_count")).map(_.value).sum + val errors = mutable.ArrayBuffer[String]() + actual match { + case Left(e) => + errors += s"Comet failed: $e\n${e.getStackTrace.take(15).mkString("\n")}" + case Right(rows) => + val exp = counts(expected) + val act = counts(rows) + if (exp != act) { + val missing = exp.toSeq.flatMap { case (k, n) => + val d = n - act.getOrElse(k, 0) + if (d > 0) Some(k -> d) else None + } + val extra = act.toSeq.flatMap { case (k, n) => + val d = n - exp.getOrElse(k, 0) + if (d > 0) Some(k -> d) else None + } + errors += s"results differ: Spark ${expected.size} rows, Comet ${rows.size} rows, " + + s"${missing.map(_._2).sum} missing in Comet, ${extra.map(_._2).sum} extra in " + + "Comet\nmissing (row x count):\n" + + missing.take(5).map { case (k, d) => s" ${showRow(k)} x$d" }.mkString("\n") + + "\nextra (row x count):\n" + + extra.take(5).map { case (k, d) => s" ${showRow(k)} x$d" }.mkString("\n") + } + } + if (native != q.expectNative) { + errors += s"expected native=${q.expectNative}, got $native" + } + if (q.expectNative && condition != q.expectCondition) { + errors += s"expected a join filter in CometSortMergeJoin=${q.expectCondition}" + } + val error = + if (errors.isEmpty) None + else Some(s"$ctx ${errors.mkString("\n")}\nquery: ${q.sql}\nplan:\n$plan") + val o = Outcome( + q.label, + shape.name, + ds.name, + expected.size, + actual.fold(_ => -1, _.size), + native, + condition, + spills, + error) + // scalastyle:off println + println(o.line) + // scalastyle:on println + o + } + } + + private def runAll(ds: DataSet, shape: Shape, qs: Seq[Q]): Seq[Outcome] = { + val outcomes = qs.map(q => runCase(ds, shape, q)) + val failed = outcomes.filter(_.error.nonEmpty) + if (failed.nonEmpty) { + fail( + s"[seed=$seed data=${ds.name} (${ds.stats}) shape=${shape.name}] " + + s"${failed.size} of ${outcomes.size} cases failed: " + + failed.map(_.label).mkString(", ") + "\n\n" + + failed.take(3).flatMap(_.error).mkString("\n\n")) + } + outcomes + } + + private def matrix( + ds: DataSet, + kinds: Seq[JoinKind], + fs: Seq[Filter], + flipped: Boolean = false, + twoKeys: Boolean = false): Seq[Q] = + for (k <- kinds; f <- fs) yield joinQuery(ds, k, f, flipped, twoKeys) + + private def assumeFull(): Unit = + assume(fullMatrix, "the full matrix runs only with -Dcomet.test.smjFuzz.full=true") + + // ---------------------------------------------------------------- default tests + + test("every join type and filter shape, batch 64") { + val ds = dataSet("main") + runAll(ds, b64, matrix(ds, joinKinds, filters)) + } + + test("flipped sides and two join keys, batch 64") { + val ds = dataSet("main") + runAll( + ds, + b64, + matrix(ds, Seq(rightOuter, leftSemi, leftAnti), Seq(interval, band), flipped = true) ++ + matrix(ds, joinKinds, Seq(interval), twoKeys = true)) + } + + test("batch 7: key groups spanning many batches") { + val ds = dataSet("main") + runAll(ds, b7, matrix(ds, joinKinds, Seq(interval, band, most, leftCast, bothCast))) + } + + test("batch 1024: key groups around and over one batch") { + val ds = dataSet("main") + runAll(ds, b1024, matrix(ds, joinKinds, Seq(interval, rare, alwaysTrue))) + } + + test("boolean and conditional operators around lifted subexpressions, batch 64") { + val ds = dataSet("main") + runAll(ds, b64, matrix(ds, joinKinds, defaultBoolFilters)) + } + + test("boolean and conditional operators, batch 7 and spilling") { + val ds = dataSet("main") + runAll(ds, b7, matrix(ds, joinKinds, Seq(caseFlag, threeValued))) + runAll(ds, spill64, matrix(ds, Seq(inner, fullOuter, leftAnti), Seq(orTs, coalesceTs))) + } + + test("spilling join under a tiny memory pool") { + val ds = dataSet("main") + val outcomes = runAll(ds, spill64, matrix(ds, joinKinds, Seq(interval, most))) + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + + test("existence joins from EXISTS and IN in a disjunction") { + val ds = dataSet("main") + runAll(ds, b64, existenceQueries(ds)) + } + + // ---------------------------------------------------------------- full matrix + + for (shape <- allShapes) { + test(s"full: ${shape.name}, every join type, filter, side order and key count") { + assumeFull() + val ds = dataSet("main") + val outcomes = runAll( + ds, + shape, + matrix(ds, joinKinds, filters) ++ + matrix(ds, joinKinds, filters, flipped = true) ++ + matrix(ds, joinKinds, filters, twoKeys = true) ++ + existenceQueries(ds)) + if (shape.spill) { + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + } + } + + for (shape <- allShapes) { + test(s"full: ${shape.name}, boolean and conditional operators, both side orders") { + assumeFull() + val ds = dataSet("main") + val outcomes = runAll( + ds, + shape, + matrix(ds, joinKinds, boolFilters) ++ matrix(ds, joinKinds, boolFilters, flipped = true)) + if (shape.spill) { + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + } + } + + for (shape <- Seq(b7, spill64)) { + test(s"full: second seed, boolean and conditional operators, ${shape.name}") { + assumeFull() + val ds = dataSet("alt") + runAll(ds, shape, matrix(ds, joinKinds, boolFilters)) + } + } + + for (shape <- Seq(b1024, spill64)) { + test(s"full: groups of thousands of rows, boolean and conditional operators, ${shape.name}") { + assumeFull() + val ds = dataSet("big") + val fs = Seq(orTs, caseFlag, caseNested, notOr, inCross, coalesceTs, threeValued) + runAll(ds, shape, matrix(ds, joinKinds, fs) ++ matrix(ds, joinKinds, fs, flipped = true)) + } + } + + for (shape <- Seq(b7, b64, spill64)) { + test(s"full: second seed, ${shape.name}") { + assumeFull() + val ds = dataSet("alt") + runAll( + ds, + shape, + matrix(ds, joinKinds, filters) ++ matrix(ds, joinKinds, filters, flipped = true)) + } + } + + for (shape <- Seq(b64, b1024, b8192, spill64)) { + test(s"full: groups of thousands of rows, ${shape.name}") { + assumeFull() + val ds = dataSet("big") + val fs = Seq(interval, band, rare, alwaysFalse, nullable, leftCast, bothCast) + runAll(ds, shape, matrix(ds, joinKinds, fs) ++ matrix(ds, joinKinds, fs, flipped = true)) + } + } +} diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 2ccc7f6f19e..25ec4c2e0eb 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -888,7 +888,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } - test("a sort-merge join adds smjCondition only with a join condition") { + test("a sort-merge join adds smjCrossCondition only with a CASE or IF over both inputs") { withTables { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", @@ -902,28 +902,39 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } val out = Width(4, 0) val equi = join("SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k") - assert(model.terms(equi, Engine.Comet) == Seq(Term(Smj, out))) - assert(model.terms(equi, Engine.Spark) == Seq(Term(Smj, out))) - same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) - val band = join( "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + "AND a.v <= b.v AND b.v <= a.v + 13") - assert(model.terms(band, Engine.Comet) == Seq(Term(Smj, out), Term(SmjCondition, out))) - assert(model.terms(band, Engine.Spark) == Seq(Term(Smj, out), Term(SmjCondition, out))) + val oneSide = join( + "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + + "AND a.v <= CASE WHEN b.v > 500 THEN b.v ELSE b.v + 13 END") + val crossCase = join(crossCaseOnT) + assert(band.originalPlan.asInstanceOf[SortMergeJoinExec].condition.isDefined) + for (engine <- Engine.all; j <- Seq(equi, band, oneSide)) { + assert(model.terms(j, engine) == Seq(Term(Smj, out))) + } + same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) val table = EngineCostTable.default for (engine <- Engine.all) { + same(model.operatorPrice(band, engine), model.operatorPrice(equi, engine)) + assert( + model.terms(crossCase, engine) == Seq(Term(Smj, out), Term(SmjCrossCondition, out))) val price: (CostClass, Width) => Double = if (engine == Engine.Comet) table.comet else table.spark - same(model.operatorPrice(band, engine), price(Smj, out) + price(SmjCondition, out)) + same( + model.operatorPrice(crossCase, engine), + price(Smj, out) + price(SmjCrossCondition, out)) } } } } + private val crossCaseOnT = + "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + + "AND CASE WHEN a.v % 2 = b.v % 2 THEN a.v - b.v ELSE 10 - (b.v - a.v) END BETWEEN 1 AND 10" + for (aqe <- Seq("false", "true")) { - test( - s"a sort-merge join with a join condition runs in Spark, without one natively (AQE=$aqe)") { + test(s"a sort-merge join runs in Spark only under a cross-input CASE or IF (AQE=$aqe)") { withTables { withAqe( aqe, @@ -935,17 +946,19 @@ class CostBasedEngineChoiceSuite extends CometTestBase { count(plan) { case j: SortMergeJoinExec => j }) val equi = "SELECT a.k, a.v, b.s FROM t a JOIN t b ON a.k = b.k" val band = equi + " AND a.v <= b.v AND b.v <= a.v + 13" - val (equiOff, equiOn) = offAndOn(run(equi)) - assert(joins(equiOff) == (1, 0), s"plan:\n$equiOff") - assert(joins(equiOn) == (1, 0), s"plan:\n$equiOn") - assert(cometOperatorNames(equiOff) == cometOperatorNames(equiOn), s"$equiOff\n$equiOn") - val (bandOff, bandOn) = offAndOn(run(band)) - assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") - assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") + for (query <- Seq(equi, band)) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (1, 0), s"plan:\n$off") + assert(joins(on) == (1, 0), s"plan:\n$on") + assert(cometOperatorNames(off) == cometOperatorNames(on), s"$off\n$on") + } + val (crossOff, crossOn) = offAndOn(run(crossCaseOnT)) + assert(joins(crossOff) == (1, 0), s"plan:\n$crossOff") + assert(joins(crossOn) == (0, 1), s"plan:\n$crossOn") withSQLConf( flag -> "true", - costTable -> "smjCondition.comet=0,0,0;smjCondition.spark=0,0") { - val plan = run(band) + costTable -> "smjCrossCondition.comet=0,0,0;smjCrossCondition.spark=0,0") { + val plan = run(crossCaseOnT) assert(joins(plan) == (1, 0), s"plan:\n$plan") } } @@ -1007,7 +1020,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "SELECT o.k, o.v, g.v AS gv FROM ev o JOIN dim p ON o.k = p.k " + "LEFT JOIN ev g ON o.k = g.k AND p.completed_dt < g.v AND o.v > g.v" - private val chargedConditions = Seq( + private val otherConditions = Seq( crossInputBounds, "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + "AND e.vs BETWEEN d.ls - 2592000 AND d.ls", @@ -1023,6 +1036,24 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l AND e.v < d.u " + "AND e.s <> d.s") + private val crossConditionals = Seq( + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k AND CASE WHEN e.s = d.s " + + "THEN e.vs - d.ls ELSE 10 - (d.ls - e.vs) END BETWEEN 1 AND 10", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l " + + "AND IF(d.u IS NULL, e.vs < d.ls + 2592000, e.v < d.u)") + + private val unchargedConditions = Seq( + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND e.v >= CASE WHEN d.s = 's1' THEN d.l ELSE d.created_at END AND e.v < d.u", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND IF(e.s = 's1', e.vs, e.vs + 3600) BETWEEN d.ls AND d.ls + 86400", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND datediff(e.v, d.l) BETWEEN 0 AND 30", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND e.vs - d.ls BETWEEN 1 AND 3600", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND (e.s = d.s OR CAST(d.ls AS string) = e.s)") + private def conditionTerms(query: String): Seq[Seq[Term]] = { val plan = runUnordered(sql(query)) val joins = nodes(plan).collect { @@ -1032,34 +1063,33 @@ class CostBasedEngineChoiceSuite extends CometTestBase { case j: SortMergeJoinExec if j.condition.isDefined => j } assert(joins.nonEmpty, s"plan:\n$plan") - joins.map(j => model.terms(j, Engine.Comet).filter(_.costClass == SmjCondition)) + joins.map(j => model.terms(j, Engine.Comet)) } for (aqe <- Seq("false", "true")) { - test( - s"a validity interval adds no smjCondition, a band or another condition does (AQE=$aqe)") { + test(s"only a CASE or IF over both inputs adds smjCrossCondition (AQE=$aqe)") { withIntervals { withAqe( aqe, flag -> "false", SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { - for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) { - val query = - s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" - assert(conditionTerms(query).forall(_.isEmpty), query) - } - for (query <- intervalQueries) { - assert(conditionTerms(query).forall(_.isEmpty), query) + val intervals = + for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) + yield s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" + for (query <- intervals ++ intervalQueries ++ otherConditions ++ unchargedConditions) { + assert(conditionTerms(query).forall(_.map(_.costClass) == Seq(Smj)), query) } - for (query <- chargedConditions) { - assert(conditionTerms(query).exists(_.nonEmpty), query) + for (query <- crossConditionals) { + assert( + conditionTerms(query).exists(_.map(_.costClass) == Seq(Smj, SmjCrossCondition)), + query) } } } } - test(s"a validity interval join stays native, a band join runs in Spark (AQE=$aqe)") { + test(s"a join under a CASE or IF over both inputs runs in Spark, others native (AQE=$aqe)") { withIntervals { withAqe( aqe, @@ -1071,17 +1101,19 @@ class CostBasedEngineChoiceSuite extends CometTestBase { count(plan) { case j: SortMergeJoinExec => j }) val interval = "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + "AND e.v >= d.l AND e.v < COALESCE(d.u, timestamp '9999-12-31')" - val (intervalOff, intervalOn) = offAndOn(run(interval)) - assert(joins(intervalOff) == (1, 0), s"plan:\n$intervalOff") - assert(joins(intervalOn) == (1, 0), s"plan:\n$intervalOn") - val (bandOff, bandOn) = offAndOn(run(chargedConditions(1))) - assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") - assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") - val (crossOff, crossOn) = offAndOn(run(crossInputBounds)) - assert(joins(crossOff) == (2, 0), s"plan:\n$crossOff") - assert( - count(crossOn) { case j: SortMergeJoinExec if j.condition.isDefined => j } == 1, - s"plan:\n$crossOn") + for ((query, native) <- Seq( + interval -> 1, + otherConditions(1) -> 1, + crossInputBounds -> 2) ++ unchargedConditions.map(_ -> 1)) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (native, 0), s"plan:\n$off") + assert(joins(on) == (native, 0), s"plan:\n$on") + } + for (query <- crossConditionals) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (1, 0), s"plan:\n$off") + assert(joins(on) == (0, 1), s"plan:\n$on") + } } } } @@ -1405,7 +1437,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { |""".stripMargin for (aqe <- Seq("false", "true")) { - test(s"fbj_order_type's band join under a root repartition runs in Spark (AQE=$aqe)") { + test(s"fbj_order_type's band join under a root repartition stays native (AQE=$aqe)") { withFbj { withAqe( aqe, @@ -1434,7 +1466,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { conditions(off) == (Seq("band", "forward", "replenishment", "replenishment"), Nil), s"plan:\n$off") val (native, spark) = conditions(on) - assert(spark.contains("band"), s"plan:\n$on") + assert(native.contains("band") && !spark.contains("band"), s"plan:\n$on") assert(native.contains("replenishment"), s"plan:\n$on") } }