diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index 3ed1120ee0..c2c0e097a3 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -335,6 +335,23 @@ Native operators reserve through DataFusion's `MemoryConsumer` / `MemoryReservat An operator that never calls `try_grow` is invisible to the pool no matter how much memory it uses. +### Sort and whole-partition windows + +The sort merge reservation is capped at 1/32 of the configured off-heap budget per +concurrent Spark task (executor cores divided by task CPUs), up to DataFusion's default. +This leaves room for input batches on small executors; the spillable merge can grow its +reservation when it needs more. It does not increase the memory pool or suppress allocation +failures. An individual batch still has to fit the available execution budget. + +`PartitionAggregateWindowExec` handles full-partition `sum`, `avg`, `count`, `min`, and +`max` frames. It updates the existing native accumulators incrementally and reserves the +retained input batches. On reservation failure it spills those rows through DataFusion's +spill manager, then replays one spill file at a time with the final aggregate columns. +Only the current window partition is retained, and small partitions avoid disk entirely. +Accumulator state is reserved separately. This preserves native execution without retaining +an entire wide partition in memory. Other window frames continue to use DataFusion's existing +window operators; this is not a general spill implementation for all window functions. + ## Crossing the FFI boundary Batches move between the JVM and native over the Arrow C Data and C Stream interfaces, which are diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index fa7e2e375c..77e6a4fb23 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -666,6 +666,7 @@ pub unsafe extern "system" fn Java_org_apache_comet_Native_createPlan( task_cpus as usize, &spark_config, &spark_plan, + (off_heap_mode != JNI_FALSE).then_some(memory_limit as usize), )?; let plan_creation_time = start.elapsed(); @@ -819,7 +820,24 @@ fn configure_skip_partial_aggregation(config: &mut SessionConfig, plan: &Operato } } +/// DataFusion's fixed 10 MiB merge reserve can consume most of a small Spark task's +/// share before the sorter admits its first batch. Cap this eager reservation at 1/32 +/// of the per-task budget. The spillable merge can grow it as needed; larger executors +/// retain the upstream default. Explicit testing overrides are applied afterwards. +fn configure_sort_spill_reservation( + config: &mut SessionConfig, + off_heap_limit: usize, + executor_cores: usize, + task_cpus: usize, +) { + let concurrent_tasks = (executor_cores / task_cpus.max(1)).max(1); + let cap = (off_heap_limit / concurrent_tasks / 32).max(1); + let reservation = &mut config.options_mut().execution.sort_spill_reservation_bytes; + *reservation = (*reservation).min(cap); +} + /// Configure DataFusion session context. +#[allow(clippy::too_many_arguments)] fn prepare_datafusion_session_context( batch_size: usize, memory_pool: Arc, @@ -828,6 +846,7 @@ fn prepare_datafusion_session_context( task_cpus: usize, spark_config: &HashMap, spark_plan: &Operator, + off_heap_limit: Option, ) -> CometResult { let paths = local_dirs.into_iter().map(PathBuf::from).collect(); let disk_manager = DiskManagerBuilder::default() @@ -845,6 +864,11 @@ fn prepare_datafusion_session_context( // modified by changing spark.task.cpus in the Spark config. .with_batch_size(batch_size); + if let Some(limit) = off_heap_limit { + let executor_cores = spark_config.get_usize(SPARK_EXECUTOR_CORES, 1); + configure_sort_spill_reservation(&mut session_config, limit, executor_cores, task_cpus); + } + // Translate the Comet-namespaced row-level pushdown flag into the equivalent // DataFusion session options. `pushdown_filters` enables the parquet reader's // RowFilter evaluation during decode (late materialization); `reorder_filters` @@ -1837,6 +1861,23 @@ mod tests { use datafusion_comet_proto::spark_expression::{AggExpr, Count, Expr, Sum}; use datafusion_comet_proto::spark_operator::{HashAggregate, ShuffleWriter}; + #[test] + fn sort_merge_reserve_scales_with_the_task_budget() { + let mut config = SessionConfig::new(); + let default = config.options().execution.sort_spill_reservation_bytes; + configure_sort_spill_reservation(&mut config, 64 * 1024 * 1024, 2, 1); + assert_eq!( + config.options().execution.sort_spill_reservation_bytes, + 1024 * 1024 + ); + let mut config = SessionConfig::new(); + configure_sort_spill_reservation(&mut config, 8 * 1024 * 1024 * 1024, 4, 1); + assert_eq!( + config.options().execution.sort_spill_reservation_bytes, + default + ); + } + #[test] fn skip_partial_eligibility_is_fail_closed() { let count = AggExpr { diff --git a/native/core/src/execution/operators/mod.rs b/native/core/src/execution/operators/mod.rs index d09b0b4fb3..a22107a382 100644 --- a/native/core/src/execution/operators/mod.rs +++ b/native/core/src/execution/operators/mod.rs @@ -42,7 +42,9 @@ pub use iceberg_write::IcebergWriteExec; mod parquet_writer; pub use parquet_writer::{ParquetCompression, ParquetWriterExec}; mod csv_scan; +mod partition_aggregate_window; pub mod projection; +pub use partition_aggregate_window::PartitionAggregateWindowExec; mod sample; pub use sample::SampleExec; mod rank_limit; diff --git a/native/core/src/execution/operators/partition_aggregate_window.rs b/native/core/src/execution/operators/partition_aggregate_window.rs new file mode 100644 index 0000000000..a5412ec010 --- /dev/null +++ b/native/core/src/execution/operators/partition_aggregate_window.rs @@ -0,0 +1,478 @@ +// 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. + +use std::collections::VecDeque; +use std::fmt::Formatter; +use std::sync::Arc; + +use arrow::array::RecordBatch; +use arrow::datatypes::SchemaRef; +use datafusion::common::tree_node::TreeNodeRecursion; +use datafusion::common::utils::evaluate_partition_ranges; +use datafusion::common::{Result, ScalarValue}; +use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; +use datafusion::execution::{SpillFile, TaskContext}; +use datafusion::logical_expr::Accumulator; +use datafusion::physical_expr::window::PlainAggregateWindowExpr; +use datafusion::physical_expr::{PhysicalExpr, PhysicalSortExpr}; +use datafusion::physical_plan::metrics::{ + BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet, SpillMetrics, +}; +use datafusion::physical_plan::spill::SpillManager; +use datafusion::physical_plan::stream::RecordBatchStreamAdapter; +use datafusion::physical_plan::windows::WindowAggExec; +use datafusion::physical_plan::{ + DisplayAs, DisplayFormatType, ExecutionPlan, InputDistributionRequirements, PlanProperties, + SendableRecordBatchStream, WindowExpr, +}; +use futures::{stream, StreamExt}; + +/// Whole-partition aggregates need the final aggregate before they can emit the first row. +/// Keep the accumulator, but spill the rows instead of buffering the entire partition in RAM. +/// Existing Spark-compatible accumulators retain their null and overflow semantics. +#[derive(Debug)] +pub struct PartitionAggregateWindowExec { + window: WindowAggExec, + metrics: ExecutionPlanMetricsSet, +} + +impl PartitionAggregateWindowExec { + pub fn supports(exprs: &[Arc]) -> bool { + !exprs.is_empty() + && exprs.iter().all(|expr| { + let frame = expr.get_window_frame(); + frame.start_bound.is_unbounded() + && frame.end_bound.is_unbounded() + && expr + .as_any() + .downcast_ref::() + .is_some_and(|agg| { + // These accumulators have bounded state; collection aggregates need + // a separate strategy for spilling their accumulator, not just rows. + matches!( + agg.get_aggregate_expr().fun().name(), + "sum" | "avg" | "count" | "min" | "max" + ) + }) + }) + } + + pub fn new(window: WindowAggExec) -> Self { + Self { + window, + metrics: ExecutionPlanMetricsSet::new(), + } + } +} + +impl DisplayAs for PartitionAggregateWindowExec { + fn fmt_as(&self, _: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { + write!(f, "PartitionAggregateWindowExec") + } +} + +impl ExecutionPlan for PartitionAggregateWindowExec { + fn name(&self) -> &str { + "PartitionAggregateWindowExec" + } + fn properties(&self) -> &Arc { + self.window.properties() + } + fn children(&self) -> Vec<&Arc> { + self.window.children() + } + fn apply_expressions( + &self, + f: &mut dyn FnMut(&Arc) -> Result, + ) -> Result { + self.window.apply_expressions(f) + } + fn maintains_input_order(&self) -> Vec { + vec![true] + } + fn input_distribution_requirements(&self) -> InputDistributionRequirements { + self.window.input_distribution_requirements() + } + fn required_input_ordering( + &self, + ) -> Vec> { + self.window.required_input_ordering() + } + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + Ok(Arc::new(Self::new(WindowAggExec::try_new( + self.window.window_expr().to_vec(), + Arc::clone(&children[0]), + !self.window.window_expr()[0].partition_by().is_empty(), + )?))) + } + fn metrics(&self) -> Option { + Some(self.metrics.clone_inner()) + } + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + let runtime = context.runtime_env(); + let input = self + .window + .input() + .execute(partition, Arc::clone(&context))?; + let schema = self.schema(); + let state = WindowState { + spill: SpillManager::new( + Arc::clone(&runtime), + SpillMetrics::new(&self.metrics, partition), + input.schema(), + ), + rows_reservation: MemoryConsumer::new("WindowRows") + .with_can_spill(true) + .register(&runtime.memory_pool), + state_reservation: MemoryConsumer::new("WindowAccumulator") + .register(&runtime.memory_pool), + baseline: BaselineMetrics::new(&self.metrics, partition), + input, + input_done: false, + schema: Arc::clone(&schema), + exprs: self.window.window_expr().to_vec(), + keys: self.window.partition_by_sort_keys()?, + pending: VecDeque::new(), + current_key: None, + accumulators: vec![], + rows: vec![], + files: VecDeque::new(), + replay: None, + result: vec![], + emitting: false, + }; + let stream = stream::try_unfold(state, |mut state| async move { + Ok(state.next_batch().await?.map(|batch| (batch, state))) + }); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream))) + } +} + +struct WindowState { + baseline: BaselineMetrics, + input: SendableRecordBatchStream, + input_done: bool, + schema: SchemaRef, + exprs: Vec>, + keys: Vec, + pending: VecDeque<(Vec, RecordBatch)>, + current_key: Option>, + accumulators: Vec>, + rows: Vec, + files: VecDeque>, + spill: SpillManager, + rows_reservation: MemoryReservation, + state_reservation: MemoryReservation, + replay: Option, + result: Vec, + emitting: bool, +} + +impl WindowState { + fn spill_rows(&mut self) -> Result<()> { + if let Some(file) = self + .spill + .spill_record_batch_and_finish(&self.rows, "window rows")? + { + self.files.push_back(file); + } + self.rows.clear(); + self.rows_reservation.free(); + Ok(()) + } + + fn append(&mut self, batch: RecordBatch) -> Result<()> { + for (expr, accumulator) in self.exprs.iter().zip(&mut self.accumulators) { + let args = expr + .expressions() + .iter() + .map(|e| e.evaluate(&batch)?.into_array(batch.num_rows())) + .collect::>>()?; + accumulator.update_batch(&args)?; + } + let state_size = self.accumulators.iter().map(|a| a.size()).sum(); + if self.state_reservation.try_resize(state_size).is_err() { + self.spill_rows()?; + self.state_reservation.try_resize(state_size)?; + } + let size = batch.get_array_memory_size(); + if self.rows_reservation.try_grow(size).is_err() { + self.spill_rows()?; + // A single input batch may itself exceed the share. Write it directly, without + // retaining it or claiming an unbounded memory reservation. + if self.rows_reservation.try_grow(size).is_err() { + self.rows.push(batch); + return self.spill_rows(); + } + } + self.rows.push(batch); + Ok(()) + } + + fn finish_partition(&mut self) -> Result<()> { + self.result = self + .accumulators + .iter_mut() + .map(|a| a.evaluate()) + .collect::>()?; + self.accumulators.clear(); + self.state_reservation.free(); + if !self.files.is_empty() { + self.spill_rows()?; + } else { + let batches = std::mem::take(&mut self.rows); + self.replay = Some(Box::pin(RecordBatchStreamAdapter::new( + self.input.schema(), + stream::iter(batches.into_iter().map(Ok)), + ))); + } + self.emitting = true; + Ok(()) + } + + async fn next_batch(&mut self) -> Result> { + loop { + if self.emitting { + if let Some(replay) = &mut self.replay { + if let Some(batch) = replay.next().await { + let batch = batch?; + let mut columns = batch.columns().to_vec(); + for value in &self.result { + columns.push(value.to_array_of_size(batch.num_rows())?); + } + self.baseline.record_output(batch.num_rows()); + return Ok(Some(RecordBatch::try_new( + Arc::clone(&self.schema), + columns, + )?)); + } + self.replay = None; + } + if let Some(file) = self.files.pop_front() { + // Open one file at a time, without prefetching the rest of the partition. + self.replay = Some(self.spill.read_spill_as_stream_unbuffered(file, None)?); + continue; + } + self.rows_reservation.free(); + self.current_key = None; + self.result.clear(); + self.emitting = false; + } + if let Some((key, batch)) = self.pending.pop_front() { + if self + .current_key + .as_ref() + .is_some_and(|current| *current != key) + { + self.pending.push_front((key, batch)); + self.finish_partition()?; + continue; + } + if self.current_key.is_none() { + self.current_key = Some(key); + self.accumulators = self + .exprs + .iter() + .map(|expr| { + expr.as_any() + .downcast_ref::() + .expect("supports checked the expression") + .get_aggregate_expr() + .create_accumulator() + }) + .collect::>()?; + } + self.append(batch)?; + continue; + } + match if self.input_done { + None + } else { + self.input.next().await + } { + Some(batch) => { + let batch = batch?; + if batch.num_rows() == 0 { + continue; + } + let keys = self + .keys + .iter() + .map(|k| k.evaluate_to_sort_column(&batch)) + .collect::>>()?; + for range in evaluate_partition_ranges(batch.num_rows(), &keys)? { + let key = keys + .iter() + .map(|k| ScalarValue::try_from_array(&k.values, range.start)) + .collect::>>()?; + self.pending + .push_back((key, batch.slice(range.start, range.end - range.start))); + } + } + None if self.current_key.is_some() => { + self.input_done = true; + self.finish_partition()?; + } + None => return Ok(None), + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{Array, Int64Array, StringArray, UInt32Array}; + use arrow::compute::SortOptions; + use arrow::datatypes::{DataType, Field, Schema}; + use datafusion::datasource::memory::MemorySourceConfig; + use datafusion::datasource::source::DataSourceExec; + use datafusion::execution::memory_pool::{GreedyMemoryPool, MemoryPool}; + use datafusion::execution::runtime_env::RuntimeEnvBuilder; + use datafusion::functions_aggregate::sum::sum_udaf; + use datafusion::logical_expr::WindowFrame; + use datafusion::physical_expr::aggregate::AggregateExprBuilder; + use datafusion::physical_expr::expressions::Column; + use datafusion::physical_expr::LexOrdering; + use datafusion::prelude::{SessionConfig, SessionContext}; + + #[tokio::test] + async fn whole_partition_aggregates_spill_and_preserve_rows() -> Result<()> { + for partitioned in [false, true] { + // Exercise no spill, buffered spill, and a batch larger than the entire budget. + for budget in [1_000_000, 16_000, 1024] { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Int64, true), + Field::new("value", DataType::Int64, true), + Field::new("payload", DataType::Utf8, false), + ])); + let keys: Vec<_> = (0..80) + .map(|i| if i < 9 { None } else { Some(i / 25) }) + .collect(); + let values: Vec<_> = (0..80) + .map(|i| if i % 3 == 0 { None } else { Some(i) }) + .collect(); + let payload = "x".repeat(1024); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int64Array::from(keys.clone())), + Arc::new(Int64Array::from(values.clone())), + Arc::new(StringArray::from(vec![payload.as_str(); 80])), + ], + )?; + let mut batches = vec![batch.slice(0, 0)]; + for start in (0..80).step_by(7) { + let indices = UInt32Array::from( + (start as u32..(start + 7).min(80) as u32).collect::>(), + ); + batches.push(arrow::compute::take_record_batch(&batch, &indices)?); + } + let key: Arc = Arc::new(Column::new("key", 0)); + let order = PhysicalSortExpr::new( + Arc::clone(&key), + SortOptions { + descending: false, + nulls_first: true, + }, + ); + let config = MemorySourceConfig::try_new(&[batches], Arc::clone(&schema), None)? + .try_with_sort_information(vec![LexOrdering::new(vec![order]).unwrap()])?; + let input = Arc::new(DataSourceExec::new(Arc::new(config))); + let aggregate = + AggregateExprBuilder::new(sum_udaf(), vec![Arc::new(Column::new("value", 1))]) + .schema(schema) + .alias("total") + .build()?; + let partition_keys = if partitioned { vec![key] } else { vec![] }; + let expr: Arc = Arc::new(PlainAggregateWindowExpr::new( + Arc::new(aggregate), + &partition_keys, + &[], + Arc::new(WindowFrame::new(None)), + None, + )); + assert!(PartitionAggregateWindowExec::supports(&[Arc::clone(&expr)])); + let plan = PartitionAggregateWindowExec::new(WindowAggExec::try_new( + vec![expr], + input, + partitioned, + )?); + let pool: Arc = Arc::new(GreedyMemoryPool::new(budget)); + let runtime = Arc::new( + RuntimeEnvBuilder::new() + .with_memory_pool(Arc::clone(&pool)) + .build()?, + ); + let ctx = SessionContext::new_with_config_rt(SessionConfig::new(), runtime); + let mut output = plan.execute(0, ctx.task_ctx())?; + let mut row = 0; + while let Some(batch) = output.next().await { + let batch = batch?; + let sums = batch + .column(3) + .as_any() + .downcast_ref::() + .unwrap(); + let payloads = batch + .column(2) + .as_any() + .downcast_ref::() + .unwrap(); + for i in 0..batch.num_rows() { + let expected: i64 = values + .iter() + .zip(&keys) + .filter(|(_, k)| !partitioned || **k == keys[row]) + .filter_map(|(v, _)| *v) + .sum(); + assert!(!sums.is_null(i)); + assert_eq!(sums.value(i), expected); + assert_eq!(payloads.value(i), payload); + assert_eq!( + ScalarValue::try_from_array(batch.column(0), i)?, + ScalarValue::Int64(keys[row]) + ); + assert_eq!( + ScalarValue::try_from_array(batch.column(1), i)?, + ScalarValue::Int64(values[row]) + ); + row += 1; + } + } + assert_eq!(row, 80); + drop(output); + assert_eq!(pool.reserved(), 0); + let spills = plan.metrics().unwrap().spill_count().unwrap_or(0); + assert_eq!(spills > 0, budget < 1_000_000); + // Dropping a partially consumed replay releases reservations and spill owners. + let mut cancelled = plan.execute(0, ctx.task_ctx())?; + assert!(cancelled.next().await.transpose()?.is_some()); + drop(cancelled); + assert_eq!(pool.reserved(), 0); + } + } + Ok(()) + } +} diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 10cc4ae37d..054cf9489d 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -43,7 +43,7 @@ use crate::execution::{ expressions::subquery::Subquery, operators::{ CometFilterExec, ExecutionError, ExpandExec, ExplodeExec, ParquetCompression, - ParquetWriterExec, SampleExec, ScanExec, ShuffleScanExec, + ParquetWriterExec, PartitionAggregateWindowExec, SampleExec, ScanExec, ShuffleScanExec, }, planner::expression_registry::ExpressionRegistry, planner::operator_registry::OperatorRegistry, @@ -2430,20 +2430,27 @@ impl PhysicalPlanner { // trigger a retract call. let window_expr = window_expr?; let all_bounded = window_expr.iter().all(|e| e.uses_bounded_memory()); - let window_agg: Arc = if all_bounded { - Arc::new(BoundedWindowAggExec::try_new( - window_expr, - Arc::clone(&child.native_plan), - InputOrderMode::Sorted, - !partition_exprs.is_empty(), - )?) - } else { - Arc::new(WindowAggExec::try_new( - window_expr, - Arc::clone(&child.native_plan), - !partition_exprs.is_empty(), - )?) - }; + let window_agg: Arc = + if PartitionAggregateWindowExec::supports(&window_expr) { + Arc::new(PartitionAggregateWindowExec::new(WindowAggExec::try_new( + window_expr, + Arc::clone(&child.native_plan), + !partition_exprs.is_empty(), + )?)) + } else if all_bounded { + Arc::new(BoundedWindowAggExec::try_new( + window_expr, + Arc::clone(&child.native_plan), + InputOrderMode::Sorted, + !partition_exprs.is_empty(), + )?) + } else { + Arc::new(WindowAggExec::try_new( + window_expr, + Arc::clone(&child.native_plan), + !partition_exprs.is_empty(), + )?) + }; // DataFusion's window functions don't always return the same Arrow // type that Spark expects (e.g. `row_number` returns UInt64 while diff --git a/native/spark-expr/src/datetime_funcs/timestamp_trunc.rs b/native/spark-expr/src/datetime_funcs/timestamp_trunc.rs index 3a40aa79a2..1d752ce96a 100644 --- a/native/spark-expr/src/datetime_funcs/timestamp_trunc.rs +++ b/native/spark-expr/src/datetime_funcs/timestamp_trunc.rs @@ -66,7 +66,13 @@ impl TimestampTruncExpr { TimestampTruncExpr { child, format, - timezone: Arc::from(timezone), + // Spark/Arrow scan and literal timestamps use UTC. Preserve the canonical Arrow + // type for its Etc/UTC alias, otherwise native comparisons reject equal timezones. + timezone: Arc::from(if timezone == "Etc/UTC" { + "UTC".to_owned() + } else { + timezone + }), } } } @@ -163,3 +169,39 @@ impl PhysicalExpr for TimestampTruncExpr { ))) } } + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::TimestampMicrosecondArray; + use arrow::datatypes::Field; + use datafusion::physical_expr::expressions::{Column, Literal}; + + #[test] + fn utc_alias_has_the_canonical_scan_timestamp_type() { + let timestamps = Arc::new( + TimestampMicrosecondArray::from(vec![Some(45_000_000_000), None]).with_timezone("UTC"), + ); + let schema = Arc::new(Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(Microsecond, Some("UTC".into())), + true, + )])); + let batch = RecordBatch::try_new(Arc::clone(&schema), vec![timestamps]).unwrap(); + for timezone in ["UTC", "Etc/UTC"] { + let expr = TimestampTruncExpr::new( + Arc::new(Column::new("ts", 0)), + Arc::new(Literal::new(Utf8(Some("day".to_owned())))), + timezone.to_owned(), + ); + assert_eq!( + expr.data_type(&schema).unwrap(), + schema.field(0).data_type().clone() + ); + let actual = expr.evaluate(&batch).unwrap().into_array(2).unwrap(); + let expected = + TimestampMicrosecondArray::from(vec![Some(0), None]).with_timezone("UTC"); + assert_eq!(actual.as_ref(), &expected); + } + } +} diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala index 83fbca6b63..232ff9e135 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala @@ -156,11 +156,9 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come // instance with a single `init(partitionIndex)` call, so `Rand` / `MonotonicallyIncreasingID` // state advances correctly across batches. // - // `ExecSubqueryExpression` (`ScalarSubquery`, `InSubqueryExec`) is accepted: the surrounding - // Comet operator's inherited `SparkPlan.waitForSubqueries` populates the subquery's - // `result` field before evaluation. The closure serializer captures that value into the - // arg-0 bytes, and the dispatcher keys its compile cache on those bytes, so distinct subquery - // results produce distinct cache entries. + // Scalar subqueries are lowered to BoundReference inputs by CometScalaUDF before this + // check. Their resolved values travel through the native subquery argument path, not + // inside the serialized kernel; the same kernel can safely serve different results. // // `Unevaluable`: rejected by default. `isCodegenInertUnevaluable` exempts version-specific // leaves that are `Unevaluable` but never invoked by codegen (e.g. Spark 4.0's @@ -372,7 +370,8 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come // leaf-only-children roots, where it is exact; see [[canShortCircuitNulls]]. val nullCheck = inputOrdinals .map(ord => - s"this.col$ord.${CometBatchKernelCodegenInput.nullCheckMethod(inputSchema(ord))}(i)") + s"this.col$ord.${CometBatchKernelCodegenInput.nullCheckMethod(inputSchema(ord))}" + + s"(i & this.col${ord}_rowMask)") .mkString(" || ") // `NullIntolerant` only constrains "any input null -> output null"; it does NOT promise // that non-null inputs always produce non-null output. `MakeTimestamp(failOnError=false)` diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala index 2ed7e33c90..642a6783bf 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -70,11 +70,17 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { classOf[IntervalMonthDayNanoVector]) private val cometPlainVectorName: String = classOf[CometPlainVector].getName + // Native scalars arrive as length-one vectors. A per-batch mask broadcasts them + // without copying their values or adding a branch to every row read. Ordinary columns + // use -1, so their row index is unchanged, including when the cached kernel is reused. + private def rowIndex(ord: Int): String = s"(this.rowIdx & this.col${ord}_rowMask)" + /** Emit kernel typed-vector field declarations for every level of every input column. */ def emitInputFieldDecls(inputSchema: Seq[ArrowColumnSpec]): String = { val lines = new mutable.ArrayBuffer[String]() inputSchema.zipWithIndex.foreach { case (spec, ord) => val path = s"col$ord" + lines += s"private int ${path}_rowMask;" collectVectorFieldDecls(path, spec, lines) } lines.mkString("\n ") @@ -87,6 +93,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { val lines = new mutable.ArrayBuffer[String]() inputSchema.zipWithIndex.foreach { case (spec, ord) => val path = s"col$ord" + lines += s"this.${path}_rowMask = inputs[$ord].getValueCount() == 1 ? 0 : -1;" collectCasts(path, spec, s"inputs[$ord]", lines) } lines.mkString("\n ") @@ -114,27 +121,27 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { case cls if wrapsInCometPlainVector(cls) => "isNullAt" case _ => "isNull" } - s" case $ord: return this.col$ord.$method(this.rowIdx);" + s" case $ord: return this.col$ord.$method(${rowIndex(ord)});" } } val booleanCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[BitVector] => - s" case $ord: return this.col$ord.getBoolean(this.rowIdx);" + s" case $ord: return this.col$ord.getBoolean(${rowIndex(ord)});" } val byteCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[TinyIntVector] => - s" case $ord: return this.col$ord.getByte(this.rowIdx);" + s" case $ord: return this.col$ord.getByte(${rowIndex(ord)});" } val shortCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[SmallIntVector] => - s" case $ord: return this.col$ord.getShort(this.rowIdx);" + s" case $ord: return this.col$ord.getShort(${rowIndex(ord)});" } val intCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[IntVector] || cls == classOf[DateDayVector] || cls == classOf[IntervalYearVector] => - s" case $ord: return this.col$ord.getInt(this.rowIdx);" + s" case $ord: return this.col$ord.getInt(${rowIndex(ord)});" } val longCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) @@ -143,27 +150,27 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { cls == classOf[TimeNanoVector] || cls == classOf[TimeStampMicroVector] || cls == classOf[TimeStampMicroTZVector] => - s" case $ord: return this.col$ord.getLong(this.rowIdx);" + s" case $ord: return this.col$ord.getLong(${rowIndex(ord)});" } val intervalCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[IntervalMonthDayNanoVector] => - s" case $ord: return this.col$ord.getInterval(this.rowIdx);" + s" case $ord: return this.col$ord.getInterval(${rowIndex(ord)});" } val floatCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[Float4Vector] => - s" case $ord: return this.col$ord.getFloat(this.rowIdx);" + s" case $ord: return this.col$ord.getFloat(${rowIndex(ord)});" } val doubleCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[Float8Vector] => - s" case $ord: return this.col$ord.getDouble(this.rowIdx);" + s" case $ord: return this.col$ord.getDouble(${rowIndex(ord)});" } val decimalCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) if cls == classOf[DecimalVector] => val known = decimalTypeByOrdinal.getOrElse(ord, None) val valueAddr = s"this.col${ord}_valueAddr" val slowField = s"this.col$ord" - val fastPath = emitDecimalFastBodyUnsafe(valueAddr, "this.rowIdx", " ") - val slowPath = emitDecimalSlowBody(slowField, "this.rowIdx", " ") + val fastPath = emitDecimalFastBodyUnsafe(valueAddr, rowIndex(ord), " ") + val slowPath = emitDecimalSlowBody(slowField, rowIndex(ord), " ") val body = known match { case Some(dt) if dt.precision <= Decimal.MAX_LONG_DIGITS => fastPath case Some(_) => slowPath @@ -184,7 +191,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { |${emitBinaryBodyUnsafe( s"this.col${ord}_valueAddr", s"this.col${ord}_offsetAddr", - "this.rowIdx", + rowIndex(ord), " ")} | }""".stripMargin } @@ -194,7 +201,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { |${emitUtf8BodyUnsafe( s"this.col${ord}_valueAddr", s"this.col${ord}_offsetAddr", - "this.rowIdx", + rowIndex(ord), " ")} | }""".stripMargin } @@ -333,7 +340,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { def emitGetArrayMethod(inputSchema: Seq[ArrowColumnSpec]): String = { val cases = inputSchema.zipWithIndex.collect { case (_: ArrayColumnSpec, ord) => s""" case $ord: { - | int __idx = this.rowIdx; + | int __idx = ${rowIndex(ord)}; | int __s = this.col$ord.getElementStartIndex(__idx); | int __e = this.col$ord.getElementEndIndex(__idx); | return new InputArray_col$ord(__s, __e - __s); @@ -359,7 +366,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { def emitGetMapMethod(inputSchema: Seq[ArrowColumnSpec]): String = { val cases = inputSchema.zipWithIndex.collect { case (_: MapColumnSpec, ord) => s""" case $ord: { - | int __idx = this.rowIdx; + | int __idx = ${rowIndex(ord)}; | int __s = this.col$ord.getElementStartIndex(__idx); | int __e = this.col$ord.getElementEndIndex(__idx); | return new InputMap_col$ord(__s, __e - __s); @@ -384,7 +391,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { /** Top-level `getStruct(int ordinal, int numFields)` switch when the schema has any struct. */ def emitGetStructMethod(inputSchema: Seq[ArrowColumnSpec]): String = { val cases = inputSchema.zipWithIndex.collect { case (_: StructColumnSpec, ord) => - s""" case $ord: return new InputStruct_col$ord(this.rowIdx);""".stripMargin + s""" case $ord: return new InputStruct_col$ord(${rowIndex(ord)});""".stripMargin } if (cases.isEmpty) { "" diff --git a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala index 9249ab280f..8fcab69cb4 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala @@ -22,7 +22,8 @@ package org.apache.comet.serde import scala.util.control.NonFatal import org.apache.spark.SparkEnv -import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSeq, BindReferences, Expression, Literal, RuntimeReplaceable, ScalaUDF} +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSeq, BindReferences, BoundReference, Expression, Literal, RuntimeReplaceable, ScalaUDF} +import org.apache.spark.sql.execution.ScalarSubquery import org.apache.spark.sql.types.BinaryType import org.apache.comet.CometConf @@ -95,7 +96,13 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { // Bind against only the AttributeReferences the tree actually reads, so ordinals align with // the data args we ship. val attrs = target.collect { case a: AttributeReference => a }.distinct - val boundExpr = BindReferences.bindReference(target, AttributeSeq(attrs)) + // Subqueries are resolved after planning. Ship their values as native arguments rather + // than capturing an unresolved ScalarSubquery in the serialized codegen closure. + val subqueries = target.collect { case s: ScalarSubquery => s }.distinct + val withSubqueryInputs = target.transform { case s: ScalarSubquery => + BoundReference(attrs.length + subqueries.indexOf(s), s.dataType, s.nullable) + } + val boundExpr = BindReferences.bindReference(withSubqueryInputs, AttributeSeq(attrs)) // Gate at plan time. Surface the reason via withFallbackReason rather than crashing Janino // at execute. @@ -143,7 +150,7 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { return None } - val dataArgs = attrs.map { a => + val dataArgs = (attrs ++ subqueries).map { a => exprToProtoInternal(a, inputs, binding).getOrElse { withFallbackReason(expr, s"$exprName: codegen dispatch: could not serialize data arg $a") return None diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala index a4365c0075..fad6ede8f2 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala @@ -115,21 +115,11 @@ case class CometNativeScanExec( if (bucketedScan) { originalPlan.outputPartitioning } else { - // Use perPartitionData.length instead of originalPlan.inputRDD.getNumPartitions. - // - // originalPlan.inputRDD triggers FileSourceScanExec's full scan pipeline including - // codegen on partition filter expressions. With DPP, this calls - // InSubqueryExec.doGenCode which requires the subquery to have finished - but - // outputPartitioning can be accessed before prepare() runs (e.g., by - // ValidateRequirements during plan validation). - // - // perPartitionData goes through serializedPartitionData, which explicitly resolves - // DPP subqueries (via updateResult()) before accessing file partitions. This is the - // same pattern CometIcebergNativeScanExec uses. - // - // This is also more correct: perPartitionData.length reflects the post-DPP partition - // count, matching what CometExecRDD actually uses in doExecuteColumnar(). - UnknownPartitioning(perPartitionData.length) + // Planning must not resolve DPP or enumerate its filtered files. AQE can inspect + // partitioning before CometPlanAdaptiveDynamicPruningFilters replaces the broadcast + // placeholders. Like FileSourceScanExec, advertise no partitioning guarantee here; + // CometExecRDD gets the actual (post-DPP) partition count at execution time. + UnknownPartitioning(0) } } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala index ed2d62a1a3..05008835ec 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala @@ -1005,6 +1005,11 @@ abstract class CometNativeExec extends CometExec { // broadcast plan. val (firstNonBroadcastPlanRDD, firstNonBroadcastPlanNumPartitions) = firstNonBroadcastPlan.get._1 match { + case plan: CometScanWithPlanData => + // File counts are execution data, not a planning-time partitioning guarantee. + // findAllPlanData above has already resolved DPP and serialized the selected files. + // Read the scan itself: the plan-data map omits scans with zero selected files. + (null.asInstanceOf[RDD[Any]], plan.perPartitionData.length) case plan: CometNativeExec => (null.asInstanceOf[RDD[Any]], plan.outputPartitioning.numPartitions) case plan => diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala index 400ec6fcd2..8a4590e1f5 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala @@ -79,10 +79,10 @@ class CometCodegenSourceSuite extends AnyFunSuite { Some("UTC")) val src = CometBatchKernelCodegen.generateSource(expr, IndexedSeq(spec)).body assert( - src.contains("if (this.col0.isNullAt(i))"), + src.contains("if (this.col0.isNullAt(i & this.col0_rowMask))"), s"expected short-circuit to use isNullAt for CometPlainVector-wrapped col0; got:\n$src") assert( - !src.contains("if (this.col0.isNull(i))"), + !src.contains("if (this.col0.isNull(i & this.col0_rowMask))"), s"expected no raw Arrow isNull on the CometPlainVector-wrapped col0; got:\n$src") } @@ -110,7 +110,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { val expr = Length(BoundReference(0, StringType, nullable = true)) val src = gen(expr, nullableString) assert( - src.contains("case 0: return this.col0.isNull(this.rowIdx);"), + src.contains("case 0: return this.col0.isNull((this.rowIdx & this.col0_rowMask));"), s"expected nullable isNullAt to delegate to the Arrow vector; got:\n$src") } @@ -125,12 +125,12 @@ class CometCodegenSourceSuite extends AnyFunSuite { test("NullIntolerant expression emits input-null short-circuit before ev.code") { // Upper is NullIntolerant (null in -> null out). Expect the default body to prepend - // `if (this.col0.isNull(i)) { setNull; } else { ... }` so null rows skip the whole + // `if (this.col0.isNull(i & this.col0_rowMask)) { setNull; } else { ... }` so null rows skip the whole // expression eval, not just the setNull write. val expr = Upper(BoundReference(0, StringType, nullable = true)) val src = gen(expr, nullableString) assert( - src.contains("this.col0.isNull(i)"), + src.contains("this.col0.isNull(i & this.col0_rowMask)"), s"expected NullIntolerant short-circuit on input ordinal 0; got:\n$src") assert( src.contains("output.setNull(i);"), @@ -144,7 +144,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { val expr = Length(Upper(BoundReference(0, StringType, nullable = true))) val src = gen(expr, nullableString) assert( - src.contains("if (this.col0.isNull(i))"), + src.contains("if (this.col0.isNull(i & this.col0_rowMask))"), s"expected short-circuit on col0 when every node is NullIntolerant; got:\n$src") } @@ -163,7 +163,8 @@ class CometCodegenSourceSuite extends AnyFunSuite { BoundReference(1, StringType, nullable = true)))) val src = gen(expr, nullable1, nullable2) assert( - !src.contains("this.col0.isNull(i) || this.col1.isNull(i)"), + !src.contains( + "this.col0.isNull(i & this.col0_rowMask) || this.col1.isNull(i & this.col1_rowMask)"), "expected no pre-null short-circuit when Concat breaks the NullIntolerant chain; " + s"got:\n$src") } @@ -190,10 +191,12 @@ class CometCodegenSourceSuite extends AnyFunSuite { BoundReference(1, IntegerType, nullable = true)) val src = gen(expr, strCol, intCol) assert( - !src.contains("this.col0.isNull(i) || this.col1.isNullAt(i)"), + !src.contains( + "this.col0.isNull(i & this.col0_rowMask) || this.col1.isNullAt(i & this.col1_rowMask)"), s"expected no union-of-inputs short-circuit when a Cast sits under the root; got:\n$src") assert( - !src.contains("if (this.col0.isNull(i))") && !src.contains("if (this.col1.isNullAt(i))"), + !src.contains("if (this.col0.isNull(i & this.col0_rowMask))") && !src.contains( + "if (this.col1.isNullAt(i & this.col1_rowMask))"), s"expected no pre-eval input-null short-circuit at all for this shape; got:\n$src") } @@ -222,7 +225,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { BoundReference(0, IntegerType, nullable = true))) val src = gen(expr, intCol) assert( - !src.contains("if (this.col0.isNullAt(i))"), + !src.contains("if (this.col0.isNullAt(i & this.col0_rowMask))"), s"expected no short-circuit when a foldable subtree under the root can raise; got:\n$src") } @@ -234,7 +237,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { Upper(Substring(BoundReference(0, StringType, nullable = true), Literal(1), Literal(2))) val src = gen(expr, nullableString) assert( - src.contains("if (this.col0.isNull(i))"), + src.contains("if (this.col0.isNull(i & this.col0_rowMask))"), s"expected the single-ordinal short-circuit to survive a Literal-only argument list; got:\n$src") } @@ -252,7 +255,8 @@ class CometCodegenSourceSuite extends AnyFunSuite { BoundReference(1, IntegerType, nullable = true)) val src = gen(expr, intCol, intCol) assert( - src.contains("if (this.col0.isNullAt(i) || this.col1.isNullAt(i))"), + src.contains( + "if (this.col0.isNullAt(i & this.col0_rowMask) || this.col1.isNullAt(i & this.col1_rowMask))"), s"expected union-of-inputs short-circuit for a leaf-only two-input tree; got:\n$src") } @@ -530,9 +534,9 @@ class CometCodegenSourceSuite extends AnyFunSuite { // The short-circuit must test every ordinal the tree reads, not just the first. assert( src.contains( - "if (this.col0.isNullAt(i) || this.col1.isNullAt(i) || " + - "this.col2.isNullAt(i) || this.col3.isNullAt(i) || this.col4.isNullAt(i) || " + - "this.col5.isNull(i))"), + "if (this.col0.isNullAt(i & this.col0_rowMask) || this.col1.isNullAt(i & this.col1_rowMask) || " + + "this.col2.isNullAt(i & this.col2_rowMask) || this.col3.isNullAt(i & this.col3_rowMask) || this.col4.isNullAt(i & this.col4_rowMask) || " + + "this.col5.isNull(i & this.col5_rowMask))"), s"expected the short-circuit to test all six ordinals; source:\n$formatted") } @@ -566,7 +570,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { "expected exactly one setNull site (the post-eval ev.isNull guard, with no short-circuit); " + s"found $setNullOccurrences. Source:\n$formatted") assert( - !src.contains("if (this.col0.isNull(i))"), + !src.contains("if (this.col0.isNull(i & this.col0_rowMask))"), s"expected no input-null short-circuit when a Cast sits under the root; source:\n$formatted") } @@ -591,7 +595,7 @@ class CometCodegenSourceSuite extends AnyFunSuite { val result = CometBatchKernelCodegen.generateSource(expr, IndexedSeq(intCol)) val src = result.body assert( - src.contains("if (this.col0.isNullAt(i))"), + src.contains("if (this.col0.isNullAt(i & this.col0_rowMask))"), s"expected input-null short-circuit for the single-input tree; got:\n$src") val setNullOccurrences = "output\\.setNull\\(i\\);".r.findAllIn(src).length assert( diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 7bb8fc0a9d..e3570e7733 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -1385,15 +1385,35 @@ class CometCodegenSuite checkSparkAnswerAndOperator(df) } + test("codegen takes unresolved scalar subqueries as runtime inputs") { + withTable("codegen_strings", "codegen_patterns") { + sql("CREATE TABLE codegen_strings (s STRING) USING parquet") + // Both rows must reach the same batch: separate one-row files hide scalar broadcasting bugs. + sql( + "INSERT INTO codegen_strings SELECT /*+ COALESCE(1) */ * " + + "FROM VALUES ('abc123'), ('def456') AS v(s)") + sql("CREATE TABLE codegen_patterns (pattern STRING) USING parquet") + sql("INSERT INTO codegen_patterns VALUES ('([0-9]+)')") + for (aggregate <- Seq("max(pattern)", "max(pattern) FILTER (WHERE false)")) { + assertCodegenRan { + checkSparkAnswerAndOperator( + sql(s"SELECT regexp_extract(s, (SELECT $aggregate FROM codegen_patterns), 1) " + + "FROM codegen_strings")) + } + } + // A second execution must use its own subquery value, even when the kernel is cached. + sql("INSERT OVERWRITE codegen_patterns VALUES ('([a-z]+)')") + assertCodegenRan { + checkSparkAnswerAndOperator( + sql("SELECT regexp_extract(s, (SELECT max(pattern) FROM codegen_patterns), 1) " + + "FROM codegen_strings")) + } + } + } + test("ScalaUDF composed with reused scalar subquery across projection and filter") { - // The same scalar subquery appears in two sites: the projection (which the dispatcher - // compiles into a fused kernel) and the filter (a separate operator). Each site holds its - // own `ScalarSubquery` expression instance with its own `@volatile result` field. Each - // surrounding operator's inherited `SparkPlan.waitForSubqueries` populates its instance's - // `result` before the dispatcher's bridge serializes the expression. The populated value - // travels through closure serialization into the cache key's bytes, so different subquery - // values compile distinct kernels. Exercises the full subquery-correctness invariant - // documented on `CometBatchKernelCodegen.canHandle`. + // Exercise reused subqueries beside a dispatched expression in separate operators. + // The preceding regression also puts a subquery inside the dispatched expression. spark.udf.register("addOne", (i: Int) => i + 1) withTable("t", "t2") { sql("CREATE TABLE t (x INT) USING parquet") diff --git a/spark/src/test/scala/org/apache/comet/CometDateTimeUtilsSuite.scala b/spark/src/test/scala/org/apache/comet/CometDateTimeUtilsSuite.scala index 072009d7ab..8ab4801218 100644 --- a/spark/src/test/scala/org/apache/comet/CometDateTimeUtilsSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometDateTimeUtilsSuite.scala @@ -35,6 +35,27 @@ class CometDateTimeUtilsSuite extends CometTestBase { import testImplicits._ + test("date_trunc UTC alias compares natively with parquet timestamps") { + withSQLConf( + "spark.sql.session.timeZone" -> "Etc/UTC", + "spark.sql.parquet.outputTimestampType" -> "TIMESTAMP_MICROS", + "spark.comet.expression.TruncTimestamp.enabled" -> "true", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false") { + withTempPath { dir => + val input = Seq("2024-01-01 12:34:56", "2024-01-02 00:00:00", null) + .toDF("value") + .selectExpr("cast(value AS TIMESTAMP) AS ts") + val df = roundtripParquet(input, dir) + checkSparkAnswerAndOperator( + df.selectExpr( + "ts", + "date_trunc('DAY', ts) AS day", + "date_trunc('DAY', ts) < ts AS earlier")) + checkSparkAnswerAndOperator(df.where("date_trunc('DAY', ts) < ts")) + } + } + } + private def roundtripParquet(df: DataFrame, tempDir: File): DataFrame = { val filename = new File(tempDir, s"dtutils_${System.currentTimeMillis()}.parquet").toString df.write.mode(SaveMode.Overwrite).parquet(filename) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/CometDppFallbackRepro3949Suite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/CometDppFallbackRepro3949Suite.scala index f672ebc082..7fa60fe1c9 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/CometDppFallbackRepro3949Suite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/CometDppFallbackRepro3949Suite.scala @@ -63,6 +63,69 @@ import org.apache.comet.{CometConf, CometExplainInfo} */ class CometDppFallbackRepro3949Suite extends CometTestBase { + test("native context preserves zero partitions after file pruning") { + withTempDir { dir => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark.range(2).selectExpr("id", "0 AS p").write.partitionBy("p").parquet(s"$dir/fact") + } + withTempView("empty_native_fact") { + spark.read.parquet(s"$dir/fact").createOrReplaceTempView("empty_native_fact") + for (adaptive <- Seq("false", "true")) { + withSQLConf( + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + "spark.sql.adaptive.enabled" -> adaptive) { + val df = sql("SELECT id + 1 FROM empty_native_fact WHERE p = 99") + val plan = unwrapAqe(df.queryExecution.executedPlan) + assert(plan.toString.contains("CometProject"), plan.toString) + val scans = plan.collect { case scan: CometNativeScanExec => scan } + assert(scans.nonEmpty, plan.toString) + assert(scans.forall(_.perPartitionData.isEmpty)) + checkAnswer(df, Seq.empty[Row]) + checkAnswer(sql("SELECT count(id) FROM empty_native_fact WHERE p = 99"), Seq(Row(0L))) + } + } + } + } + } + + test("AQE coalescing a union sibling must not execute native scan DPP during planning") { + withTempDir { dir => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(4) + .selectExpr("id AS fact_id", "id % 2 AS fact_key") + .write + .partitionBy("fact_key") + .parquet(s"$dir/fact") + spark.range(2).selectExpr("id AS dim_id", "id AS dim_key").write.parquet(s"$dir/dim") + } + withTempView("aqe_fact", "aqe_dim") { + spark.read.parquet(s"$dir/fact").createOrReplaceTempView("aqe_fact") + spark.read.parquet(s"$dir/dim").createOrReplaceTempView("aqe_dim") + withSQLConf( + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false", + "spark.comet.exec.union.enabled" -> "false", + "spark.sql.adaptive.enabled" -> "true", + "spark.sql.adaptive.coalescePartitions.enabled" -> "true", + "spark.sql.shuffle.partitions" -> "4", + "spark.sql.optimizer.dynamicPartitionPruning.enabled" -> "true") { + val df = sql(""" + SELECT /*+ BROADCAST(d) */ cast(f.fact_id AS STRING) AS x + FROM aqe_fact f JOIN aqe_dim d ON f.fact_key = d.dim_key + WHERE d.dim_id < 1 + UNION ALL + SELECT cast(sum(fact_id) AS STRING) FROM aqe_fact GROUP BY fact_key + """) + val initial = unwrapAqe(df.queryExecution.executedPlan) + assert(initial.toString.contains("CometNativeScan"), initial.toString) + assert(initial.toString.contains("dynamicpruning"), initial.toString) + checkAnswer(df, Seq(Row("0"), Row("2"), Row("2"), Row("4"))) + } + } + } + } + // ---------------------------------------------------------------------- // Mechanism (synthetic): proves the AQE wrap flips the fallback decision. // ----------------------------------------------------------------------