From 47ce539e4d6881180b3ee4e63eb2632998bd9f33 Mon Sep 17 00:00:00 2001 From: msaf Date: Fri, 9 Oct 2026 13:06:27 +0100 Subject: [PATCH] perf: stream simple window functions over sorted input without per-partition slicing Add CometSortedWindowExec for windows whose expressions are all ROW_NUMBER, RANK, DENSE_RANK, or LEAD/LAG with a constant offset and default and without IGNORE NULLS. Partition and peer boundaries are computed per batch with vectorized comparisons, window columns are produced for the whole batch, and input columns pass through untouched, so tiny window partitions no longer pay for slicing every column, per-partition evaluator state and concatenation. Other windows keep BoundedWindowAggExec / PartitionAggregateWindowExec / WindowAggExec. spark.comet.exec.window.sorted.enabled (default true) switches it off. CometWindowExec now shows the native output rows and compute time. Co-Authored-By: Claude Opus 5.5 (1M context) --- native/core/Cargo.toml | 4 + native/core/benches/sorted_window.rs | 251 ++++ native/core/src/execution/jni_api.rs | 10 +- native/core/src/execution/operators/mod.rs | 4 + .../src/execution/operators/sorted_window.rs | 1302 +++++++++++++++++ native/core/src/execution/planner.rs | 242 ++- native/core/src/execution/spark_config.rs | 1 + .../scala/org/apache/comet/CometConf.scala | 12 + .../org/apache/comet/CometExecIterator.scala | 1 + .../spark/sql/comet/CometWindowExec.scala | 3 +- .../comet/exec/CometWindowExecSuite.scala | 152 ++ 11 files changed, 1976 insertions(+), 6 deletions(-) create mode 100644 native/core/benches/sorted_window.rs create mode 100644 native/core/src/execution/operators/sorted_window.rs diff --git a/native/core/Cargo.toml b/native/core/Cargo.toml index fc5631ce7c6..ef48269d3af 100644 --- a/native/core/Cargo.toml +++ b/native/core/Cargo.toml @@ -157,3 +157,7 @@ harness = false [[bench]] name = "sort_payload" harness = false + +[[bench]] +name = "sorted_window" +harness = false diff --git a/native/core/benches/sorted_window.rs b/native/core/benches/sorted_window.rs new file mode 100644 index 00000000000..b684a20cbbc --- /dev/null +++ b/native/core/benches/sorted_window.rs @@ -0,0 +1,251 @@ +// 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. + +//! `SortedWindowExec` against DataFusion's `BoundedWindowAggExec` on sorted input with tiny +//! window partitions, for `LEAD` and `ROW_NUMBER`, over narrow and wide nested rows. + +use std::sync::Arc; + +use arrow::array::{ArrayRef, Int64Array, ListArray, StringArray, StructArray}; +use arrow::buffer::OffsetBuffer; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, FieldRef, Fields, Schema, SchemaRef}; +use arrow::record_batch::RecordBatch; +use comet::execution::operators::{SortedWindowExec, SortedWindowFunction}; +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; +use datafusion::common::ScalarValue; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::datasource::source::DataSourceExec; +use datafusion::functions_window::lead_lag::lead_udwf; +use datafusion::functions_window::row_number::row_number_udwf; +use datafusion::logical_expr::{WindowFrame, WindowFunctionDefinition}; +use datafusion::physical_expr::expressions::{Column, Literal}; +use datafusion::physical_expr::{LexOrdering, PhysicalExpr, PhysicalSortExpr}; +use datafusion::physical_plan::windows::{create_window_expr, BoundedWindowAggExec}; +use datafusion::physical_plan::{collect, ExecutionPlan, InputOrderMode}; +use datafusion::prelude::SessionContext; +use tokio::runtime::Runtime; + +const ROWS_PER_BATCH: usize = 8192; +const BATCHES: usize = 8; +const WIDE_COLUMNS: usize = 8; + +#[derive(Clone, Copy)] +enum Function { + Lead, + RowNumber, +} + +fn schema(wide: bool) -> SchemaRef { + let mut fields = vec![ + Field::new("key", DataType::Int64, false), + Field::new("ts", DataType::Int64, false), + ]; + if wide { + for i in 0..WIDE_COLUMNS { + fields.push(Field::new(format!("s{i}"), DataType::Utf8, true)); + fields.push(Field::new( + format!("n{i}"), + DataType::Struct(nested_fields()), + true, + )); + } + } + Arc::new(Schema::new(fields)) +} + +fn nested_fields() -> Fields { + Fields::from(vec![ + Field::new("a", DataType::Int64, true), + Field::new_list("b", Field::new_list_field(DataType::Utf8, true), true), + ]) +} + +fn batches(sizes: &[usize], wide: bool) -> Vec { + let total = ROWS_PER_BATCH * BATCHES; + let mut keys = Vec::with_capacity(total); + let mut key = 0i64; + let mut i = 0; + while keys.len() < total { + for _ in 0..sizes[i % sizes.len()] { + keys.push(key); + } + key += 1; + i += 1; + } + keys.truncate(total); + let schema = schema(wide); + keys.chunks(ROWS_PER_BATCH) + .map(|chunk| { + let n = chunk.len(); + let mut columns: Vec = vec![ + Arc::new(Int64Array::from(chunk.to_vec())), + Arc::new(Int64Array::from_iter_values((0..n as i64).map(|v| v * 7))), + ]; + if wide { + for c in 0..WIDE_COLUMNS { + columns.push(Arc::new(StringArray::from_iter_values( + (0..n).map(|r| format!("value-{c}-{r}-padding")), + ))); + let strings = + StringArray::from_iter_values((0..n * 3).map(|r| format!("element-{r}"))); + let list = ListArray::new( + Arc::new(Field::new_list_field(DataType::Utf8, true)), + OffsetBuffer::from_lengths(std::iter::repeat_n(3, n)), + Arc::new(strings), + None, + ); + columns.push(Arc::new(StructArray::new( + nested_fields(), + vec![ + Arc::new(Int64Array::from_iter_values(0..n as i64)), + Arc::new(list), + ], + None, + ))); + } + } + RecordBatch::try_new(Arc::clone(&schema), columns).unwrap() + }) + .collect() +} + +fn input(batches: &[RecordBatch], schema: &SchemaRef) -> Arc { + let ordering = LexOrdering::new(vec![ + PhysicalSortExpr { + expr: Arc::new(Column::new("key", 0)), + options: SortOptions::default(), + }, + PhysicalSortExpr { + expr: Arc::new(Column::new("ts", 1)), + options: SortOptions::default(), + }, + ]) + .unwrap(); + let config = MemorySourceConfig::try_new(&[batches.to_vec()], Arc::clone(schema), None) + .unwrap() + .try_with_sort_information(vec![ordering]) + .unwrap(); + Arc::new(DataSourceExec::new(Arc::new(config))) +} + +fn plan( + batches: &[RecordBatch], + schema: &SchemaRef, + function: Function, + sorted: bool, +) -> Arc { + let partition_by: Vec> = vec![Arc::new(Column::new("key", 0))]; + let order_by = vec![PhysicalSortExpr { + expr: Arc::new(Column::new("ts", 1)), + options: SortOptions::default(), + }]; + let ts: Arc = Arc::new(Column::new("ts", 1)); + let default = ScalarValue::Int64(Some(i64::MAX)); + let (def, name, args) = match function { + Function::Lead => ( + lead_udwf(), + "lead", + vec![ + Arc::clone(&ts), + Arc::new(Literal::new(ScalarValue::Int64(Some(1)))) as Arc, + Arc::new(Literal::new(default.clone())), + ], + ), + Function::RowNumber => (row_number_udwf(), "row_number", vec![]), + }; + let window_expr = create_window_expr( + &WindowFunctionDefinition::WindowUDF(def), + name.to_string(), + &args, + &partition_by, + &order_by, + Arc::new(WindowFrame::new(Some(true))), + Arc::clone(schema), + false, + false, + None, + ) + .unwrap(); + let input = input(batches, schema); + if !sorted { + return Arc::new( + BoundedWindowAggExec::try_new(vec![window_expr], input, InputOrderMode::Sorted, true) + .unwrap(), + ); + } + let (function, field): (SortedWindowFunction, FieldRef) = match function { + Function::Lead => ( + SortedWindowFunction::Shift { + value: ts, + offset: 1, + default, + }, + window_expr.field().unwrap(), + ), + Function::RowNumber => ( + SortedWindowFunction::RowNumber, + Arc::new( + window_expr + .field() + .unwrap() + .as_ref() + .clone() + .with_data_type(DataType::Int32), + ), + ), + }; + Arc::new( + SortedWindowExec::try_new(input, partition_by, order_by, vec![function], vec![field]) + .unwrap(), + ) +} + +fn bench(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("sorted_window"); + group.sample_size(10); + for (shape, sizes) in [("1row", vec![1usize]), ("2.2rows", vec![2, 2, 2, 3, 2])] { + for wide in [false, true] { + let schema = schema(wide); + let data = batches(&sizes, wide); + for (function, fname) in [ + (Function::Lead, "lead"), + (Function::RowNumber, "row_number"), + ] { + for sorted in [false, true] { + let id = format!( + "{fname}/{shape}/{}/{}", + if wide { "wide" } else { "narrow" }, + if sorted { "sorted" } else { "bounded" } + ); + group.bench_function(BenchmarkId::from_parameter(id), |b| { + b.iter(|| { + let plan = plan(&data, &schema, function, sorted); + rt.block_on(collect(plan, SessionContext::new().task_ctx())) + .unwrap() + }) + }); + } + } + } + } + group.finish(); +} + +criterion_group!(benches, bench); +criterion_main!(benches); diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index 19a2d0dd75d..8fcaf964ca4 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -107,7 +107,9 @@ use tokio::sync::mpsc; use tokio::task::JoinHandle; use crate::execution::memory_pools::{create_memory_pool, parse_memory_pool_config}; -use crate::execution::operators::{PartitionAggregateWindowEnabled, ScanExec, ShuffleScanExec}; +use crate::execution::operators::{ + PartitionAggregateWindowEnabled, ScanExec, ShuffleScanExec, SortedWindowEnabled, +}; use crate::execution::shuffle::{ decode_remote_shuffle_batch, read_ipc_compressed, CompressionCodec, ShuffleReadCoalescer, ShuffleWriterExec, @@ -122,7 +124,7 @@ use crate::execution::memory_pools::logging_pool::LoggingMemoryPool; use crate::execution::spark_config::{ SparkConfig, COMET_DEBUG_ENABLED, COMET_DEBUG_MEMORY, COMET_EXEC_SORT_SPILL_BEFORE_OUTPUT_THRESHOLD, COMET_EXEC_WINDOW_PARTITION_AGGREGATE_ENABLED, - COMET_EXPLAIN_NATIVE_ENABLED, COMET_MAX_TEMP_DIRECTORY_SIZE, + COMET_EXEC_WINDOW_SORTED_ENABLED, COMET_EXPLAIN_NATIVE_ENABLED, COMET_MAX_TEMP_DIRECTORY_SIZE, COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED, COMET_TRACING_ENABLED, SPARK_EXECUTOR_CORES, }; use crate::parquet::encryption_support::{CometEncryptionFactory, ENCRYPTION_FACTORY_ID}; @@ -970,6 +972,10 @@ fn prepare_datafusion_session_context( session_config = session_config.with_extension(Arc::new(PartitionAggregateWindowEnabled)); } + if spark_config.get_bool(COMET_EXEC_WINDOW_SORTED_ENABLED) { + session_config = session_config.with_extension(Arc::new(SortedWindowEnabled)); + } + configure_skip_partial_aggregation(&mut session_config, spark_plan); let runtime = rt_config.build()?; diff --git a/native/core/src/execution/operators/mod.rs b/native/core/src/execution/operators/mod.rs index 76826d8e2a7..b651b26f9b5 100644 --- a/native/core/src/execution/operators/mod.rs +++ b/native/core/src/execution/operators/mod.rs @@ -55,8 +55,12 @@ mod rank_limit; pub use rank_limit::{PartitionedRankLimitExec, WindowFnKind}; mod scan; mod shuffle_scan; +mod sorted_window; pub use csv_scan::init_csv_datasource_exec; pub use shuffle_scan::ShuffleScanExec; +pub use sorted_window::{ + sorted_window_supports_output_type, SortedWindowEnabled, SortedWindowExec, SortedWindowFunction, +}; /// Fixtures for the nested-nullability drift from /// , shared by the `expand` and diff --git a/native/core/src/execution/operators/sorted_window.rs b/native/core/src/execution/operators/sorted_window.rs new file mode 100644 index 00000000000..d82898bd416 --- /dev/null +++ b/native/core/src/execution/operators/sorted_window.rs @@ -0,0 +1,1302 @@ +// 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. + +//! Streaming window operator for input sorted by `[partition_by..., order_by...]`, limited to +//! `ROW_NUMBER`, `RANK`, `DENSE_RANK` and `LEAD` / `LAG` with a constant offset and default +//! (without `IGNORE NULLS`). +//! +//! Every input batch is processed as a whole: partition and peer boundaries come from +//! vectorized comparisons of adjacent rows (the first row against the last row of the previous +//! batch), the window columns are computed for the whole batch, and the input columns pass +//! through untouched. Only running counters, the last keys and, for `LAG`, the last values +//! cross batch boundaries. A batch with a `LEAD` waits until enough following rows have +//! arrived to resolve its last rows. + +use std::collections::VecDeque; +use std::fmt::Formatter; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use arrow::array::{ + make_comparator, new_empty_array, Array, ArrayRef, BooleanBufferBuilder, Int32Array, + Int64Array, RecordBatch, RecordBatchOptions, UInt64Array, +}; +use arrow::buffer::BooleanBuffer; +use arrow::compute::kernels::cmp::distinct; +use arrow::compute::{concat, interleave, SortOptions}; +use arrow::datatypes::{DataType, FieldRef, Schema, SchemaRef}; +use datafusion::common::tree_node::TreeNodeRecursion; +use datafusion::common::{internal_err, Result, ScalarValue}; +use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; +use datafusion::execution::TaskContext; +use datafusion::physical_expr::{ + EquivalenceProperties, LexOrdering, OrderingRequirements, PhysicalExpr, PhysicalSortExpr, +}; +use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType}; +use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet}; +use datafusion::physical_plan::{ + apply_expression_roots, DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, + PlanProperties, RecordBatchStream, SendableRecordBatchStream, +}; +use futures::{Stream, StreamExt}; + +pub const SORTED_WINDOW_MAX_OFFSET: i64 = 1024; + +#[derive(Debug)] +pub struct SortedWindowEnabled; + +#[derive(Debug, Clone)] +pub enum SortedWindowFunction { + RowNumber, + Rank, + DenseRank, + Shift { + value: Arc, + offset: i64, + default: ScalarValue, + }, +} + +impl SortedWindowFunction { + fn offset(&self) -> i64 { + match self { + SortedWindowFunction::Shift { offset, .. } => *offset, + _ => 0, + } + } + + fn name(&self) -> String { + match self { + SortedWindowFunction::RowNumber => "row_number".to_string(), + SortedWindowFunction::Rank => "rank".to_string(), + SortedWindowFunction::DenseRank => "dense_rank".to_string(), + SortedWindowFunction::Shift { value, offset, .. } if *offset >= 0 => { + format!("lead({value}, {offset})") + } + SortedWindowFunction::Shift { value, offset, .. } => { + format!("lag({value}, {})", offset.unsigned_abs()) + } + } + } +} + +pub fn sorted_window_supports_output_type( + function: &SortedWindowFunction, + data_type: &DataType, + input_schema: &Schema, +) -> Result { + Ok(match function { + SortedWindowFunction::Shift { + value, + offset, + default, + } => { + offset.unsigned_abs() <= SORTED_WINDOW_MAX_OFFSET as u64 + && &value.data_type(input_schema)? == data_type + && &default.data_type() == data_type + && !data_type.is_null() + } + _ => matches!( + data_type, + DataType::Int32 | DataType::Int64 | DataType::UInt64 + ), + }) +} + +#[derive(Debug)] +pub struct SortedWindowExec { + input: Arc, + partition_by: Vec>, + order_by: Vec, + functions: Vec, + fields: Vec, + schema: SchemaRef, + cache: Arc, + metrics: ExecutionPlanMetricsSet, +} + +impl SortedWindowExec { + pub fn try_new( + input: Arc, + partition_by: Vec>, + order_by: Vec, + functions: Vec, + fields: Vec, + ) -> Result { + if functions.len() != fields.len() { + return internal_err!("SortedWindowExec needs one output field per function"); + } + let input_schema = input.schema(); + for (function, field) in functions.iter().zip(fields.iter()) { + if !sorted_window_supports_output_type(function, field.data_type(), &input_schema)? { + return internal_err!( + "SortedWindowExec does not support {} with output type {}", + function.name(), + field.data_type() + ); + } + } + let mut schema_fields: Vec = input_schema.fields().iter().cloned().collect(); + schema_fields.extend(fields.iter().cloned()); + let schema = Arc::new(Schema::new_with_metadata( + schema_fields, + input_schema.metadata().clone(), + )); + let mut eq_properties = EquivalenceProperties::new(Arc::clone(&schema)); + if let Some(ordering) = input.output_ordering() { + eq_properties.add_ordering(ordering.iter().cloned()); + } + let cache = Arc::new(PlanProperties::new( + eq_properties, + input.output_partitioning().clone(), + EmissionType::Incremental, + Boundedness::Bounded, + )); + Ok(Self { + input, + partition_by, + order_by, + functions, + fields, + schema, + cache, + metrics: ExecutionPlanMetricsSet::new(), + }) + } + + fn required_ordering(&self) -> Option { + let sort_exprs: Vec = self + .partition_by + .iter() + .map(|e| PhysicalSortExpr { + expr: Arc::clone(e), + options: SortOptions::default(), + }) + .chain(self.order_by.iter().cloned()) + .collect(); + LexOrdering::new(sort_exprs) + } +} + +impl DisplayAs for SortedWindowExec { + fn fmt_as(&self, t: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { + match t { + DisplayFormatType::Default | DisplayFormatType::Verbose => { + let functions = self + .functions + .iter() + .map(|e| e.name()) + .collect::>() + .join(", "); + let partition = self + .partition_by + .iter() + .map(|e| e.to_string()) + .collect::>() + .join(", "); + let order = self + .order_by + .iter() + .map(|e| e.to_string()) + .collect::>() + .join(", "); + write!( + f, + "CometSortedWindowExec: functions=[{functions}], partition_by=[{partition}], order_by=[{order}]" + ) + } + DisplayFormatType::TreeRender => write!(f, ""), + } + } +} + +impl ExecutionPlan for SortedWindowExec { + fn name(&self) -> &str { + "CometSortedWindowExec" + } + + fn properties(&self) -> &Arc { + &self.cache + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.input] + } + + fn apply_expressions( + &self, + f: &mut dyn FnMut(&Arc) -> Result, + ) -> Result { + let values = self.functions.iter().filter_map(|e| match e { + SortedWindowFunction::Shift { value, .. } => Some(value), + _ => None, + }); + apply_expression_roots( + self.partition_by + .iter() + .chain(self.order_by.iter().map(|e| &e.expr)) + .chain(values), + f, + ) + } + + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + if children.len() != 1 { + return internal_err!("SortedWindowExec takes exactly one child"); + } + Ok(Arc::new(SortedWindowExec::try_new( + Arc::clone(&children[0]), + self.partition_by.clone(), + self.order_by.clone(), + self.functions.clone(), + self.fields.clone(), + )?)) + } + + fn required_input_ordering(&self) -> Vec> { + vec![self.required_ordering().map(OrderingRequirements::from)] + } + + fn maintains_input_order(&self) -> Vec { + vec![true] + } + + fn metrics(&self) -> Option { + Some(self.metrics.clone_inner()) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + let input = self.input.execute(partition, Arc::clone(&context))?; + let needs_peers = self.functions.iter().any(|f| { + matches!( + f, + SortedWindowFunction::Rank | SortedWindowFunction::DenseRank + ) + }); + let max_lead = self + .functions + .iter() + .map(|f| f.offset().max(0) as usize) + .max() + .unwrap_or(0); + let max_lag = self + .functions + .iter() + .map(|f| (-f.offset()).max(0) as usize) + .max() + .unwrap_or(0); + let mut defaults = Vec::with_capacity(self.functions.len()); + let mut context_values = Vec::with_capacity(self.functions.len()); + for (function, field) in self.functions.iter().zip(self.fields.iter()) { + match function { + SortedWindowFunction::Shift { default, .. } => { + defaults.push(Some(default.to_array_of_size(1)?)); + context_values.push(Some(new_empty_array(field.data_type()))); + } + _ => { + defaults.push(None); + context_values.push(None); + } + } + } + let reservation = MemoryConsumer::new(format!("SortedWindowExec[{partition}]")) + .register(context.memory_pool()); + Ok(Box::pin(SortedWindowStream { + input, + schema: Arc::clone(&self.schema), + partition_by: self.partition_by.clone(), + order_by: self.order_by.iter().map(|e| Arc::clone(&e.expr)).collect(), + functions: self.functions.clone(), + output_types: self.fields.iter().map(|f| f.data_type().clone()).collect(), + needs_peers, + max_lead, + max_lag, + defaults, + seen_rows: false, + last_partition_key: None, + last_order_key: None, + row_number: 0, + rank: 0, + dense_rank: 0, + pending: VecDeque::new(), + pending_rows: 0, + context_values, + context_starts: Vec::new(), + finished: false, + reservation, + baseline_metrics: BaselineMetrics::new(&self.metrics, partition), + })) + } +} + +struct PendingBatch { + batch: RecordBatch, + starts: BooleanBuffer, + columns: Vec>, + values: Vec>, +} + +struct SortedWindowStream { + input: SendableRecordBatchStream, + schema: SchemaRef, + partition_by: Vec>, + order_by: Vec>, + functions: Vec, + output_types: Vec, + needs_peers: bool, + max_lead: usize, + max_lag: usize, + defaults: Vec>, + seen_rows: bool, + last_partition_key: Option>, + last_order_key: Option>, + row_number: u64, + rank: u64, + dense_rank: u64, + pending: VecDeque, + pending_rows: usize, + context_values: Vec>, + context_starts: Vec, + finished: bool, + reservation: MemoryReservation, + baseline_metrics: BaselineMetrics, +} + +fn supports_distinct(data_type: &DataType) -> bool { + let leaf = match data_type { + DataType::Dictionary(_, v) => v.as_ref(), + dt => dt, + }; + !leaf.is_nested() + && !matches!( + leaf, + DataType::Dictionary(_, _) | DataType::RunEndEncoded(_, _) + ) +} + +fn adjacent_changes(column: &ArrayRef) -> Result { + let len = column.len() - 1; + let previous = column.slice(0, len); + let current = column.slice(1, len); + if supports_distinct(column.data_type()) { + return Ok(distinct(&previous, ¤t)?.values().clone()); + } + let cmp = make_comparator(previous.as_ref(), current.as_ref(), SortOptions::default())?; + Ok((0..len).map(|i| !cmp(i, i).is_eq()).collect()) +} + +fn group_starts( + columns: &[ArrayRef], + last: Option<&[ArrayRef]>, + num_rows: usize, +) -> Result> { + let mut acc: Option = None; + for (idx, column) in columns.iter().enumerate() { + let first = match last { + None => true, + Some(last) => { + let cmp = + make_comparator(last[idx].as_ref(), column.as_ref(), SortOptions::default())?; + !cmp(0, 0).is_eq() + } + }; + let mut builder = BooleanBufferBuilder::new(num_rows); + builder.append(first); + if num_rows > 1 { + builder.append_buffer(&adjacent_changes(column)?); + } + let starts = builder.finish(); + acc = Some(match acc { + None => starts, + Some(acc) => &acc | &starts, + }); + } + Ok(acc) +} + +fn evaluate_columns(exprs: &[Arc], batch: &RecordBatch) -> Result> { + let num_rows = batch.num_rows(); + exprs + .iter() + .map(|e| e.evaluate(batch).and_then(|v| v.into_array(num_rows))) + .collect() +} + +fn last_row(columns: &[ArrayRef]) -> Vec { + columns.iter().map(|c| c.slice(c.len() - 1, 1)).collect() +} + +fn counter_array(values: &[u64], data_type: &DataType) -> Result { + Ok(match data_type { + DataType::Int32 => Arc::new(Int32Array::from_iter_values( + values.iter().map(|v| *v as i32), + )), + DataType::Int64 => Arc::new(Int64Array::from_iter_values( + values.iter().map(|v| *v as i64), + )), + DataType::UInt64 => Arc::new(UInt64Array::from_iter_values(values.iter().copied())), + other => return internal_err!("SortedWindowExec cannot produce {other} counters"), + }) +} + +impl SortedWindowStream { + fn ingest(&mut self, batch: RecordBatch) -> Result<()> { + let num_rows = batch.num_rows(); + if num_rows == 0 { + return Ok(()); + } + + let partition_columns = evaluate_columns(&self.partition_by, &batch)?; + let starts = match group_starts( + &partition_columns, + self.last_partition_key.as_deref(), + num_rows, + )? { + Some(starts) => starts, + None => { + let mut builder = BooleanBufferBuilder::new(num_rows); + builder.append(!self.seen_rows); + builder.append_n(num_rows - 1, false); + builder.finish() + } + }; + self.seen_rows = true; + if !partition_columns.is_empty() { + self.last_partition_key = Some(last_row(&partition_columns)); + } + + let peers = if self.needs_peers && !self.order_by.is_empty() { + let order_columns = evaluate_columns(&self.order_by, &batch)?; + let order_starts = + group_starts(&order_columns, self.last_order_key.as_deref(), num_rows)?; + self.last_order_key = Some(last_row(&order_columns)); + order_starts.map(|o| &o | &starts) + } else { + None + }; + + let wants = |kind: fn(&SortedWindowFunction) -> bool| self.functions.iter().any(kind); + let wants_row_number = wants(|f| matches!(f, SortedWindowFunction::RowNumber)); + let wants_rank = wants(|f| matches!(f, SortedWindowFunction::Rank)); + let wants_dense_rank = wants(|f| matches!(f, SortedWindowFunction::DenseRank)); + let mut row_numbers = Vec::with_capacity(if wants_row_number { num_rows } else { 0 }); + let mut ranks = Vec::with_capacity(if wants_rank { num_rows } else { 0 }); + let mut dense_ranks = Vec::with_capacity(if wants_dense_rank { num_rows } else { 0 }); + let mut row_number = self.row_number; + let mut rank = self.rank; + let mut dense_rank = self.dense_rank; + for i in 0..num_rows { + if starts.value(i) { + row_number = 0; + dense_rank = 0; + } + row_number += 1; + if self.needs_peers && (starts.value(i) || peers.as_ref().is_some_and(|p| p.value(i))) { + rank = row_number; + dense_rank += 1; + } + if wants_row_number { + row_numbers.push(row_number); + } + if wants_rank { + ranks.push(rank); + } + if wants_dense_rank { + dense_ranks.push(dense_rank); + } + } + self.row_number = row_number; + self.rank = rank; + self.dense_rank = dense_rank; + + let mut columns = Vec::with_capacity(self.functions.len()); + let mut values = Vec::with_capacity(self.functions.len()); + for (idx, function) in self.functions.iter().enumerate() { + let counters = match function { + SortedWindowFunction::Shift { value, .. } => { + columns.push(None); + values.push(Some(value.evaluate(&batch)?.into_array(num_rows)?)); + continue; + } + SortedWindowFunction::RowNumber => &row_numbers, + SortedWindowFunction::Rank => &ranks, + SortedWindowFunction::DenseRank => &dense_ranks, + }; + columns.push(Some(counter_array(counters, &self.output_types[idx])?)); + values.push(None); + } + + self.pending_rows += num_rows; + self.pending.push_back(PendingBatch { + batch, + starts, + columns, + values, + }); + self.update_reservation(); + Ok(()) + } + + fn update_reservation(&mut self) { + let size: usize = self + .pending + .iter() + .map(|p| p.batch.get_array_memory_size()) + .sum(); + self.reservation.resize(size); + } + + fn front_ready(&self) -> bool { + match self.pending.front() { + None => false, + Some(front) => { + self.finished + || self.max_lead == 0 + || self.pending_rows - front.batch.num_rows() >= self.max_lead + } + } + } + + fn emit_front(&mut self) -> Result { + let front = match self.pending.pop_front() { + Some(front) => front, + None => return internal_err!("SortedWindowExec has no pending batch"), + }; + let num_rows = front.batch.num_rows(); + self.pending_rows -= num_rows; + + let context_rows = self.context_starts.len(); + let mut starts: Vec = Vec::with_capacity(context_rows + num_rows + self.max_lead); + starts.extend_from_slice(&self.context_starts); + starts.extend(front.starts.iter()); + let mut ahead: Vec<(usize, &PendingBatch)> = Vec::new(); + let mut ahead_rows = 0; + for next in self.pending.iter() { + if ahead_rows >= self.max_lead { + break; + } + let take = (self.max_lead - ahead_rows).min(next.batch.num_rows()); + starts.extend(next.starts.iter().take(take)); + ahead.push((take, next)); + ahead_rows += take; + } + let total = starts.len(); + let mut partition_ids: Vec = Vec::with_capacity(total); + let mut current: u32 = 0; + for (idx, start) in starts.iter().enumerate() { + if idx > 0 && *start { + current += 1; + } + partition_ids.push(current); + } + + let mut columns: Vec = Vec::with_capacity(self.schema.fields().len()); + columns.extend(front.batch.columns().iter().cloned()); + for (idx, function) in self.functions.iter().enumerate() { + match function { + SortedWindowFunction::Shift { offset, .. } => { + let context = self.context_values[idx].as_ref().unwrap(); + let current_values = front.values[idx].as_ref().unwrap(); + let default = self.defaults[idx].as_ref().unwrap(); + let mut arrays: Vec<&dyn Array> = Vec::with_capacity(ahead.len() + 3); + arrays.push(context.as_ref()); + arrays.push(current_values.as_ref()); + let mut ahead_ends: Vec = Vec::with_capacity(ahead.len()); + let mut end = context_rows + num_rows; + for (take, next) in ahead.iter() { + arrays.push(next.values[idx].as_ref().unwrap().as_ref()); + end += take; + ahead_ends.push(end); + } + let default_idx = arrays.len(); + arrays.push(default.as_ref()); + let mut indices: Vec<(usize, usize)> = Vec::with_capacity(num_rows); + for i in 0..num_rows { + let row = context_rows + i; + let target = row as i64 + offset; + if target < 0 + || target >= total as i64 + || partition_ids[target as usize] != partition_ids[row] + { + indices.push((default_idx, 0)); + continue; + } + let target = target as usize; + if target < context_rows { + indices.push((0, target)); + } else if target < context_rows + num_rows { + indices.push((1, target - context_rows)); + } else { + let mut begin = context_rows + num_rows; + for (segment, segment_end) in ahead_ends.iter().enumerate() { + if target < *segment_end { + indices.push((segment + 2, target - begin)); + break; + } + begin = *segment_end; + } + } + } + columns.push(interleave(&arrays, &indices)?); + } + _ => columns.push(Arc::clone(front.columns[idx].as_ref().unwrap())), + } + } + + if self.max_lag > 0 { + let keep = self.max_lag.min(context_rows + num_rows); + for (idx, value) in front.values.iter().enumerate() { + if let Some(value) = value { + let context = self.context_values[idx].as_ref().unwrap(); + let combined = if keep <= num_rows { + value.slice(num_rows - keep, keep) + } else { + let from_context = keep - num_rows; + concat(&[ + context + .slice(context_rows - from_context, from_context) + .as_ref(), + value.as_ref(), + ])? + }; + self.context_values[idx] = Some(combined); + } + } + self.context_starts = + starts[context_rows + num_rows - keep..context_rows + num_rows].to_vec(); + } + + self.update_reservation(); + let options = RecordBatchOptions::new().with_row_count(Some(num_rows)); + Ok(RecordBatch::try_new_with_options( + Arc::clone(&self.schema), + columns, + &options, + )?) + } +} + +impl Stream for SortedWindowStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + loop { + if self.front_ready() { + let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); + let emitted = { + let _timer = elapsed_compute.timer(); + self.emit_front() + }; + return match emitted { + Ok(batch) => self + .baseline_metrics + .record_poll(Poll::Ready(Some(Ok(batch)))), + Err(e) => Poll::Ready(Some(Err(e))), + }; + } + if self.finished { + self.reservation.free(); + return Poll::Ready(None); + } + match self.input.poll_next_unpin(cx) { + Poll::Ready(Some(Ok(batch))) => { + let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); + let ingested = { + let _timer = elapsed_compute.timer(); + self.ingest(batch) + }; + if let Err(e) = ingested { + return Poll::Ready(Some(Err(e))); + } + } + Poll::Ready(Some(Err(e))) => return Poll::Ready(Some(Err(e))), + Poll::Ready(None) => self.finished = true, + Poll::Pending => return Poll::Pending, + } + } + } +} + +impl RecordBatchStream for SortedWindowStream { + fn schema(&self) -> SchemaRef { + Arc::clone(&self.schema) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{Int32Builder, Int64Builder, ListBuilder, StringBuilder, StructArray}; + use arrow::compute::{cast, concat_batches}; + use arrow::datatypes::{Field, Fields}; + use datafusion::datasource::memory::MemorySourceConfig; + use datafusion::datasource::source::DataSourceExec; + use datafusion::functions_window::lead_lag::{lag_udwf, lead_udwf}; + use datafusion::functions_window::rank::{dense_rank_udwf, rank_udwf}; + use datafusion::functions_window::row_number::row_number_udwf; + use datafusion::logical_expr::{WindowFrame, WindowFunctionDefinition}; + use datafusion::physical_expr::expressions::{Column, Literal}; + use datafusion::physical_expr::window::WindowExpr; + use datafusion::physical_plan::collect; + use datafusion::physical_plan::windows::{create_window_expr, BoundedWindowAggExec}; + use datafusion::physical_plan::InputOrderMode; + use datafusion::prelude::SessionContext; + use rand::rngs::StdRng; + use rand::{RngExt, SeedableRng}; + + #[derive(Clone, Copy)] + enum Sizes { + Ones, + Twos, + Mixed, + Geometric, + Huge, + } + + #[derive(Clone)] + enum Func { + RowNumber, + Rank, + DenseRank, + Lead(&'static str, i64, ScalarValue), + Lag(&'static str, i64, ScalarValue), + } + + fn struct_fields() -> Fields { + Fields::from(vec![ + Field::new("a", DataType::Int32, true), + Field::new_list("b", Field::new_list_field(DataType::Int64, true), true), + ]) + } + + fn schema() -> SchemaRef { + Arc::new(Schema::new(vec![ + Field::new("k1", DataType::Int32, true), + Field::new("k2", DataType::Utf8, true), + Field::new("st", DataType::Struct(struct_fields()), true), + Field::new("o", DataType::Int64, true), + Field::new("v", DataType::Int64, true), + Field::new("s", DataType::Utf8, true), + ])) + } + + fn partition_sizes(sizes: Sizes, rng: &mut StdRng) -> Vec { + match sizes { + Sizes::Ones => vec![1; 700], + Sizes::Twos => vec![2; 350], + Sizes::Mixed => (0..320).map(|_| rng.random_range(1..=4)).collect(), + Sizes::Geometric => { + let mut out = vec![]; + while out.iter().sum::() < 900 { + let mut n = 1; + while n < 60 && rng.random_bool(0.8) { + n += 1; + } + out.push(n); + } + out + } + Sizes::Huge => { + let mut out = vec![1, 2, 1]; + out.push(2500); + out.extend([1, 3, 1]); + out + } + } + } + + fn generate(sizes: Sizes, seed: u64) -> RecordBatch { + let mut rng = StdRng::seed_from_u64(seed); + let sizes = partition_sizes(sizes, &mut rng); + let mut k1 = Int32Builder::new(); + let mut k2 = StringBuilder::new(); + let mut st_a = Int32Builder::new(); + let mut st_b = ListBuilder::new(Int64Builder::new()); + let mut st_valid = vec![]; + let mut o = Int64Builder::new(); + let mut v = Int64Builder::new(); + let mut s = StringBuilder::new(); + for (p, size) in sizes.iter().enumerate() { + let tie = rng.random_range(1..=3); + for j in 0..*size { + if p < 2 { + k1.append_null(); + } else { + k1.append_value((p / 2) as i32); + } + if p % 2 == 0 { + k2.append_null(); + } else { + k2.append_value(format!("k{p}")); + } + if p == 0 { + st_a.append_null(); + st_b.append_null(); + st_valid.push(false); + } else { + if p % 3 == 0 { + st_a.append_null(); + } else { + st_a.append_value((p % 7) as i32); + } + st_b.values().append_value(p as i64); + if p % 4 == 0 { + st_b.values().append_null(); + } + st_b.append(true); + st_valid.push(true); + } + if j < 1 && rng.random_bool(0.2) { + o.append_null(); + } else { + o.append_value((j / tie) as i64); + } + if rng.random_bool(0.2) { + v.append_null(); + } else { + v.append_value(rng.random_range(-1000..1000)); + } + if rng.random_bool(0.2) { + s.append_null(); + } else { + s.append_value(format!("s{}", rng.random_range(0..100))); + } + } + } + let st = StructArray::new( + struct_fields(), + vec![ + Arc::new(st_a.finish()) as ArrayRef, + Arc::new(st_b.finish()) as ArrayRef, + ], + Some(st_valid.into()), + ); + RecordBatch::try_new( + schema(), + vec![ + Arc::new(k1.finish()), + Arc::new(k2.finish()), + Arc::new(st), + Arc::new(o.finish()), + Arc::new(v.finish()), + Arc::new(s.finish()), + ], + ) + .unwrap() + } + + fn split(batch: &RecordBatch, sizes: &[usize], seed: u64) -> Vec { + let mut rng = StdRng::seed_from_u64(seed); + let mut out = vec![]; + let mut offset = 0; + let mut i = 0; + while offset < batch.num_rows() { + let size = if sizes.is_empty() { + rng.random_range(1..=40) + } else { + sizes[i % sizes.len()] + }; + let len = size.min(batch.num_rows() - offset); + out.push(batch.slice(offset, len)); + offset += len; + i += 1; + } + out + } + + fn col(name: &str) -> Arc { + let schema = schema(); + Arc::new(Column::new(name, schema.index_of(name).unwrap())) + } + + fn lit(value: ScalarValue) -> Arc { + Arc::new(Literal::new(value)) + } + + fn order_by(options: SortOptions) -> Vec { + vec![PhysicalSortExpr { + expr: col("o"), + options, + }] + } + + fn df_window_expr( + func: &Func, + partition_by: &[Arc], + order: &[PhysicalSortExpr], + ) -> Arc { + let (def, name, args) = match func { + Func::RowNumber => (row_number_udwf(), "row_number", vec![]), + Func::Rank => (rank_udwf(), "rank", vec![]), + Func::DenseRank => (dense_rank_udwf(), "dense_rank", vec![]), + Func::Lead(c, n, d) => ( + lead_udwf(), + "lead", + vec![col(c), lit(ScalarValue::Int64(Some(*n))), lit(d.clone())], + ), + Func::Lag(c, n, d) => ( + lag_udwf(), + "lag", + vec![col(c), lit(ScalarValue::Int64(Some(*n))), lit(d.clone())], + ), + }; + create_window_expr( + &WindowFunctionDefinition::WindowUDF(def), + name.to_string(), + &args, + partition_by, + order, + Arc::new(WindowFrame::new(Some(true))), + schema(), + false, + false, + None, + ) + .unwrap() + } + + fn sorted_function(func: &Func) -> SortedWindowFunction { + match func { + Func::RowNumber => SortedWindowFunction::RowNumber, + Func::Rank => SortedWindowFunction::Rank, + Func::DenseRank => SortedWindowFunction::DenseRank, + Func::Lead(c, n, d) | Func::Lag(c, n, d) => { + let value = col(c); + let value_type = value.data_type(&schema()).unwrap(); + let default = if d.is_null() { + ScalarValue::try_from(&value_type).unwrap() + } else { + d.cast_to(&value_type).unwrap() + }; + SortedWindowFunction::Shift { + value, + offset: if matches!(func, Func::Lead(..)) { + *n + } else { + -*n + }, + default, + } + } + } + } + + async fn run_both( + batches: Vec, + funcs: &[Func], + partition_by: &[Arc], + options: SortOptions, + counter_type: DataType, + ) -> (RecordBatch, RecordBatch, usize) { + let order = order_by(options); + let input_batches = batches.len(); + let df_exprs: Vec<_> = funcs + .iter() + .map(|f| df_window_expr(f, partition_by, &order)) + .collect(); + let ordering: Vec = partition_by + .iter() + .map(|e| PhysicalSortExpr { + expr: Arc::clone(e), + options: SortOptions::default(), + }) + .chain(order.iter().cloned()) + .collect(); + let config = MemorySourceConfig::try_new(&[batches], schema(), None) + .unwrap() + .try_with_sort_information(vec![LexOrdering::new(ordering).unwrap()]) + .unwrap(); + let input: Arc = Arc::new(DataSourceExec::new(Arc::new(config))); + let expected = Arc::new( + BoundedWindowAggExec::try_new( + df_exprs.clone(), + Arc::clone(&input) as Arc, + InputOrderMode::Sorted, + !partition_by.is_empty(), + ) + .unwrap(), + ); + let functions: Vec<_> = funcs.iter().map(sorted_function).collect(); + let fields: Vec = df_exprs + .iter() + .zip(functions.iter()) + .map(|(e, f)| { + let field = e.field().unwrap().as_ref().clone(); + let field = match f { + SortedWindowFunction::Shift { .. } => field, + _ => field + .with_data_type(counter_type.clone()) + .with_nullable(false), + }; + Arc::new(field) + }) + .collect(); + let actual = Arc::new( + SortedWindowExec::try_new(input, partition_by.to_vec(), order, functions, fields) + .unwrap(), + ); + let ctx = SessionContext::new().task_ctx(); + let expected_batches = collect(expected, Arc::clone(&ctx)).await.unwrap(); + let actual_batches = collect(Arc::clone(&actual) as _, ctx).await.unwrap(); + assert_eq!(actual_batches.len(), input_batches); + let metrics = actual.metrics().unwrap(); + assert_eq!( + metrics.output_rows().unwrap(), + actual_batches.iter().map(|b| b.num_rows()).sum::() + ); + let expected = concat_batches(&expected_batches[0].schema(), &expected_batches).unwrap(); + let actual = concat_batches(&actual.schema(), &actual_batches).unwrap(); + (expected, actual, input_batches) + } + + fn assert_same(expected: &RecordBatch, actual: &RecordBatch, context: &str) { + assert_eq!(expected.num_columns(), actual.num_columns(), "{context}"); + assert_eq!(expected.num_rows(), actual.num_rows(), "{context}"); + for idx in 0..expected.num_columns() { + let e = expected.column(idx); + let a = cast(actual.column(idx), e.data_type()).unwrap(); + if e.as_ref() != a.as_ref() { + for row in 0..e.len() { + let ev = ScalarValue::try_from_array(e, row).unwrap(); + let av = ScalarValue::try_from_array(&a, row).unwrap(); + assert_eq!(ev, av, "{context}: column {idx} row {row}"); + } + } + } + } + + fn all_functions() -> Vec { + vec![ + Func::RowNumber, + Func::Rank, + Func::DenseRank, + Func::Lead("v", 1, ScalarValue::Null), + Func::Lead("v", 2, ScalarValue::Int64(Some(42))), + Func::Lead("v", 3, ScalarValue::Int32(Some(-7))), + Func::Lag("v", 1, ScalarValue::Null), + Func::Lag("v", 3, ScalarValue::Int64(Some(-1))), + Func::Lead("s", 1, ScalarValue::Utf8(Some("x".to_string()))), + Func::Lag("s", 2, ScalarValue::Null), + Func::Lead("st", 1, ScalarValue::Null), + Func::Lag("o", 1, ScalarValue::Int64(Some(0))), + Func::Lead("o", 0, ScalarValue::Null), + ] + } + + #[tokio::test] + async fn matches_bounded_window_agg_exec() { + let function_sets: Vec> = vec![ + all_functions(), + vec![Func::Lead("o", 1, ScalarValue::Int64(Some(i64::MAX)))], + vec![Func::RowNumber], + vec![Func::Rank, Func::DenseRank], + vec![Func::Lag("v", 2, ScalarValue::Null)], + ]; + let partition_bys: Vec>> = vec![ + vec![col("k1")], + vec![col("k1"), col("k2")], + vec![col("st")], + vec![], + ]; + let splits: Vec> = vec![ + vec![1], + vec![2], + vec![3], + vec![5, 1, 2], + vec![64], + vec![4096], + vec![], + ]; + let options = [ + SortOptions::default(), + SortOptions { + descending: true, + nulls_first: false, + }, + ]; + let mut seed = 0; + for sizes in [ + Sizes::Ones, + Sizes::Twos, + Sizes::Mixed, + Sizes::Geometric, + Sizes::Huge, + ] { + for (f, funcs) in function_sets.iter().enumerate() { + for (p, partition_by) in partition_bys.iter().enumerate() { + for (s, split_sizes) in splits.iter().enumerate() { + seed += 1; + let data = generate(sizes, seed); + let batches = split(&data, split_sizes, seed); + let option = options[seed as usize % 2]; + let counter_type = if seed % 2 == 0 { + DataType::UInt64 + } else { + DataType::Int32 + }; + let (expected, actual, _) = + run_both(batches, funcs, partition_by, option, counter_type).await; + assert_same( + &expected, + &actual, + &format!("seed {seed} functions {f} partition_by {p} split {s}"), + ); + } + } + } + } + } + + #[tokio::test] + async fn empty_batches_and_empty_input() { + let data = generate(Sizes::Mixed, 7); + let mut batches = vec![data.slice(0, 0)]; + for b in split(&data, &[3], 7) { + batches.push(b); + batches.push(data.slice(0, 0)); + } + let non_empty = batches.iter().filter(|b| b.num_rows() > 0).count(); + let order = order_by(SortOptions::default()); + let funcs = all_functions(); + let df_exprs: Vec<_> = funcs + .iter() + .map(|f| df_window_expr(f, &[col("k1")], &order)) + .collect(); + let fields: Vec = df_exprs + .iter() + .zip(funcs.iter()) + .map(|(e, f)| match f { + Func::Lead(..) | Func::Lag(..) => e.field().unwrap(), + _ => Arc::new( + e.field() + .unwrap() + .as_ref() + .clone() + .with_data_type(DataType::Int32) + .with_nullable(false), + ), + }) + .collect(); + let input = MemorySourceConfig::try_new_exec(&[batches], schema(), None).unwrap(); + let plan = Arc::new( + SortedWindowExec::try_new( + input, + vec![col("k1")], + order.clone(), + funcs.iter().map(sorted_function).collect(), + fields.clone(), + ) + .unwrap(), + ); + let out = collect(plan, SessionContext::new().task_ctx()) + .await + .unwrap(); + assert_eq!(out.len(), non_empty); + assert_eq!( + out.iter().map(|b| b.num_rows()).sum::(), + data.num_rows() + ); + + let input = MemorySourceConfig::try_new_exec(&[vec![]], schema(), None).unwrap(); + let plan = Arc::new( + SortedWindowExec::try_new( + input, + vec![col("k1")], + order, + funcs.iter().map(sorted_function).collect(), + fields, + ) + .unwrap(), + ); + let out = collect(plan, SessionContext::new().task_ctx()) + .await + .unwrap(); + assert!(out.is_empty()); + } + + #[tokio::test] + async fn input_columns_pass_through() { + let data = generate(Sizes::Twos, 3); + let batches = split(&data, &[100], 3); + let (_, actual, _) = run_both( + batches, + &[Func::Lead("v", 1, ScalarValue::Null), Func::RowNumber], + &[col("k1"), col("k2")], + SortOptions::default(), + DataType::Int32, + ) + .await; + for idx in 0..data.num_columns() { + assert_eq!(actual.column(idx).as_ref(), data.column(idx).as_ref()); + } + let lead = actual + .column(data.num_columns()) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(lead.len(), data.num_rows()); + } + + #[test] + fn rejects_unsupported_output_types() { + let schema = schema(); + assert!(!sorted_window_supports_output_type( + &SortedWindowFunction::RowNumber, + &DataType::Utf8, + &schema + ) + .unwrap()); + assert!(!sorted_window_supports_output_type( + &SortedWindowFunction::Shift { + value: col("v"), + offset: SORTED_WINDOW_MAX_OFFSET + 1, + default: ScalarValue::Int64(None), + }, + &DataType::Int64, + &schema + ) + .unwrap()); + assert!(!sorted_window_supports_output_type( + &SortedWindowFunction::Shift { + value: col("v"), + offset: 1, + default: ScalarValue::Int32(None), + }, + &DataType::Int64, + &schema + ) + .unwrap()); + assert!(sorted_window_supports_output_type( + &SortedWindowFunction::Shift { + value: col("v"), + offset: -SORTED_WINDOW_MAX_OFFSET, + default: ScalarValue::Int64(None), + }, + &DataType::Int64, + &schema + ) + .unwrap()); + } +} diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 620fb99b8f2..1febe577337 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -40,14 +40,17 @@ use crate::execution::operators::AlignedArrowStreamReader; use crate::execution::operators::DynamicFilterJoinExec; use crate::execution::operators::IcebergScanExec; use crate::execution::operators::IcebergWriteExec; -use crate::execution::operators::{PartitionedRankLimitExec, WindowFnKind}; +use crate::execution::operators::{ + sorted_window_supports_output_type, PartitionedRankLimitExec, WindowFnKind, +}; use crate::execution::{ expressions::list_positions::ListPositionsExpr, expressions::subquery::Subquery, operators::{ CometFilterExec, ExecutionError, ExpandExec, ExplodeExec, ParquetCompression, ParquetWriterExec, PartitionAggregateWindowEnabled, PartitionAggregateWindowExec, - SampleExec, ScanExec, ShuffleScanExec, + SampleExec, ScanExec, ShuffleScanExec, SortedWindowEnabled, SortedWindowExec, + SortedWindowFunction, }, planner::expression_registry::ExpressionRegistry, planner::operator_registry::OperatorRegistry, @@ -2519,6 +2522,23 @@ impl PhysicalPlanner { .copied_config() .get_extension::() .is_some(); + let sorted_window = if self + .session_ctx + .copied_config() + .get_extension::() + .is_some() + { + self.plan_sorted_window( + wnd, + &window_expr, + &child.native_plan, + partition_exprs, + sort_exprs, + &input_schema, + )? + } else { + None + }; let ignore_nulls = wnd.window_expr.iter().map(|e| e.ignore_nulls).collect(); let partition_aggregate = if !all_bounded && partition_aggregate_enabled { PartitionAggregateWindowExec::try_plan( @@ -2530,7 +2550,9 @@ impl PhysicalPlanner { } else { None }; - let window_agg: Arc = if all_bounded { + let window_agg: Arc = if let Some(plan) = sorted_window { + plan + } else if all_bounded { Arc::new(BoundedWindowAggExec::try_new( window_expr, Arc::clone(&child.native_plan), @@ -3548,6 +3570,109 @@ impl PhysicalPlanner { .map_err(|e| ExecutionError::DataFusionError(e.to_string())) } + fn plan_sorted_window( + &self, + wnd: &spark_operator::Window, + window_expr: &[Arc], + input: &Arc, + partition_exprs: &[Arc], + sort_exprs: &[PhysicalSortExpr], + input_schema: &SchemaRef, + ) -> Result>, ExecutionError> { + let mut functions = Vec::with_capacity(window_expr.len()); + let mut fields = Vec::with_capacity(window_expr.len()); + for (spark_expr, df_expr) in wnd.window_expr.iter().zip(window_expr.iter()) { + if spark_expr.agg_func.is_some() { + return Ok(None); + } + let Some(ExprStruct::ScalarFunc(func)) = spark_expr + .built_in_window_function + .as_ref() + .and_then(|f| f.expr_struct.as_ref()) + else { + return Ok(None); + }; + let function = match func.func.as_str() { + "row_number" => SortedWindowFunction::RowNumber, + "rank" => SortedWindowFunction::Rank, + "dense_rank" => SortedWindowFunction::DenseRank, + name @ ("lead" | "lag") => { + if spark_expr.ignore_nulls || func.args.len() != 3 { + return Ok(None); + } + let value = self.create_expr(&func.args[0], Arc::clone(input_schema))?; + let offset = self.create_expr(&func.args[1], Arc::clone(input_schema))?; + let default = self.create_expr(&func.args[2], Arc::clone(input_schema))?; + let (Some(offset), Some(default)) = ( + offset.downcast_ref::(), + default.downcast_ref::(), + ) else { + return Ok(None); + }; + let offset = match offset.value() { + ScalarValue::Int8(Some(v)) => *v as i64, + ScalarValue::Int16(Some(v)) => *v as i64, + ScalarValue::Int32(Some(v)) => *v as i64, + ScalarValue::Int64(Some(v)) => *v, + _ => return Ok(None), + }; + let Some(offset) = (if name == "lead" { + Some(offset) + } else { + offset.checked_neg() + }) else { + return Ok(None); + }; + let value_type = value.data_type(input_schema)?; + let default = if default.value().is_null() { + ScalarValue::try_from(&value_type)? + } else { + match default.value().cast_to(&value_type) { + Ok(v) => v, + Err(_) => return Ok(None), + } + }; + SortedWindowFunction::Shift { + value, + offset, + default, + } + } + _ => return Ok(None), + }; + let result_type = spark_expr.result_type.as_ref().map(to_arrow_datatype); + let data_type = match &function { + SortedWindowFunction::Shift { value, .. } => { + let value_type = value.data_type(input_schema)?; + if result_type.as_ref().is_some_and(|t| t != &value_type) { + return Ok(None); + } + value_type + } + _ => result_type.unwrap_or(DataType::UInt64), + }; + if !sorted_window_supports_output_type(&function, &data_type, input_schema)? { + return Ok(None); + } + let nullable = matches!(function, SortedWindowFunction::Shift { .. }); + let field = df_expr + .field()? + .as_ref() + .clone() + .with_data_type(data_type) + .with_nullable(nullable); + fields.push(Arc::new(field)); + functions.push(function); + } + Ok(Some(Arc::new(SortedWindowExec::try_new( + Arc::clone(input), + partition_exprs.to_vec(), + sort_exprs.to_vec(), + functions, + fields, + )?))) + } + fn process_agg_func( &self, agg_func: &AggExpr, @@ -6508,6 +6633,117 @@ mod tests { assert_eq!("ScanExec", projection_exec.children[0].native_plan.name()); } + fn window_operator(functions: Vec<(&str, Vec, bool)>) -> Operator { + let int_literal = |value: Option| Expr { + expr_struct: Some(Literal(spark_expression::Literal { + value: value.map(literal::Value::IntVal), + datatype: Some(create_proto_datatype()), + is_null: value.is_none(), + })), + ..Default::default() + }; + let order = Expr { + expr_struct: Some(SortOrder(Box::new(spark_expression::SortOrder { + child: Some(Box::new(create_bound_reference(0))), + direction: 0, + null_ordering: 0, + }))), + ..Default::default() + }; + let frame = spark_operator::WindowFrame { + frame_type: spark_operator::WindowFrameType::Rows as i32, + lower_bound: Some(spark_operator::LowerWindowFrameBound { + lower_frame_bound_struct: Some( + spark_operator::lower_window_frame_bound::LowerFrameBoundStruct::UnboundedPreceding( + spark_operator::UnboundedPreceding {}, + ), + ), + }), + upper_bound: Some(spark_operator::UpperWindowFrameBound { + upper_frame_bound_struct: Some( + spark_operator::upper_window_frame_bound::UpperFrameBoundStruct::CurrentRow( + spark_operator::CurrentRow {}, + ), + ), + }), + }; + let window_expr = functions + .into_iter() + .map(|(func, args, ignore_nulls)| { + let args = if func == "lead" || func == "lag" { + let mut args = args; + args.insert(1, int_literal(Some(1))); + args.push(int_literal(None)); + args + } else { + args + }; + spark_operator::WindowExpr { + built_in_window_function: Some(Expr { + expr_struct: Some(ScalarFunc(spark_expression::ScalarFunc { + func: func.to_string(), + args, + return_type: None, + fail_on_error: false, + })), + ..Default::default() + }), + agg_func: None, + spec: Some(spark_operator::WindowSpecDefinition { + partition_spec: vec![create_bound_reference(0)], + order_spec: vec![order.clone()], + frame_specification: Some(frame.clone()), + }), + ignore_nulls, + result_type: Some(create_proto_datatype()), + } + }) + .collect(); + Operator { + plan_id: 0, + sql_text_pool: vec![], + children: vec![create_scan()], + op_struct: Some(OpStruct::Window(Box::new(spark_operator::Window { + window_expr, + order_by_list: vec![order], + partition_by_list: vec![create_bound_reference(0)], + child: None, + }))), + } + } + + #[test] + fn sorted_window_routing() { + use crate::execution::operators::SortedWindowEnabled; + let plan_name = |enabled: bool, op: &Operator| { + let config = if enabled { + SessionConfig::new().with_extension(Arc::new(SortedWindowEnabled)) + } else { + SessionConfig::new() + }; + let planner = + PhysicalPlanner::new(Arc::new(SessionContext::new_with_config(config)), 0); + let (_, _, plan) = planner.create_plan(op, &mut vec![], 1).unwrap(); + plan.native_plan.name().to_string() + }; + let simple = window_operator(vec![ + ("row_number", vec![], false), + ("rank", vec![], false), + ("dense_rank", vec![], false), + ("lead", vec![create_bound_reference(0)], false), + ("lag", vec![create_bound_reference(0)], false), + ]); + assert_eq!(plan_name(true, &simple), "CometSortedWindowExec"); + assert_ne!(plan_name(false, &simple), "CometSortedWindowExec"); + let ignore_nulls = window_operator(vec![ + ("row_number", vec![], false), + ("lead", vec![create_bound_reference(0)], true), + ]); + assert_ne!(plan_name(true, &ignore_nulls), "CometSortedWindowExec"); + let percent_rank = window_operator(vec![("percent_rank", vec![], false)]); + assert_ne!(plan_name(true, &percent_rank), "CometSortedWindowExec"); + } + fn create_bound_reference(index: i32) -> Expr { Expr { expr_struct: Some(Bound(spark_expression::BoundReference { diff --git a/native/core/src/execution/spark_config.rs b/native/core/src/execution/spark_config.rs index 140b65609f9..805765fc2a6 100644 --- a/native/core/src/execution/spark_config.rs +++ b/native/core/src/execution/spark_config.rs @@ -28,6 +28,7 @@ pub(crate) const COMET_EXEC_SORT_SPILL_BEFORE_OUTPUT_THRESHOLD: &str = "spark.comet.exec.sort.spillBeforeOutputThreshold"; pub(crate) const COMET_EXEC_WINDOW_PARTITION_AGGREGATE_ENABLED: &str = "spark.comet.exec.window.partitionAggregate.enabled"; +pub(crate) const COMET_EXEC_WINDOW_SORTED_ENABLED: &str = "spark.comet.exec.window.sorted.enabled"; pub(crate) const SPARK_EXECUTOR_CORES: &str = "spark.executor.cores"; /// Comet configs read through this trait must be resolved by the JVM first: diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 627be899e40..03bd2701754 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -785,6 +785,18 @@ object CometConf extends ShimCometConf { .checkValue(_ >= 1, "Must be >= 1.") .createWithDefault(50) + val COMET_EXEC_WINDOW_SORTED_ENABLED: ConfigEntry[Boolean] = + conf(s"$COMET_EXEC_CONFIG_PREFIX.window.sorted.enabled") + .category(CATEGORY_EXEC) + .doc( + "Whether native window operators whose expressions are all ROW_NUMBER, RANK, " + + "DENSE_RANK, or LEAD/LAG with a constant offset and default and without IGNORE " + + "NULLS run in Comet's SortedWindowExec, which processes each sorted input batch as " + + "a whole instead of slicing it per window partition. When false, they run in " + + "DataFusion's BoundedWindowAggExec, as in upstream Comet.") + .booleanConf + .createWithDefault(true) + val COMET_EXEC_WINDOW_PARTITION_AGGREGATE_ENABLED: ConfigEntry[Boolean] = conf(s"$COMET_EXEC_CONFIG_PREFIX.window.partitionAggregate.enabled") .category(CATEGORY_EXEC) diff --git a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala index af825520317..b071de9e3f4 100644 --- a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala +++ b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala @@ -678,6 +678,7 @@ object CometExecIterator extends Logging { CometConf.COMET_DEBUG_ENABLED, CometConf.COMET_DEBUG_MEMORY_ENABLED, CometConf.COMET_EXEC_WINDOW_PARTITION_AGGREGATE_ENABLED, + CometConf.COMET_EXEC_WINDOW_SORTED_ENABLED, CometConf.COMET_EXPLAIN_NATIVE_ENABLED, CometConf.COMET_MAX_TEMP_DIRECTORY_SIZE, CometConf.COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala index 410402d114e..70ad396729e 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala @@ -656,7 +656,8 @@ case class CometWindowExec( override lazy val metrics: Map[String, SQLMetric] = Map( "dataSize" -> SQLMetrics.createSizeMetric(sparkContext, "data size"), - "numPartitions" -> SQLMetrics.createMetric(sparkContext, "number of partitions")) + "numPartitions" -> SQLMetrics.createMetric(sparkContext, "number of partitions")) ++ + CometMetricNode.baselineMetrics(sparkContext) override def outputOrdering: Seq[SortOrder] = child.outputOrdering diff --git a/spark/src/test/scala/org/apache/comet/exec/CometWindowExecSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometWindowExecSuite.scala index 4b334b34fb0..9a092bd4980 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometWindowExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometWindowExecSuite.scala @@ -19,6 +19,8 @@ package org.apache.comet.exec +import java.sql.Timestamp + import scala.util.Random import org.scalactic.source.Position @@ -1673,4 +1675,154 @@ class CometWindowExecSuite extends CometTestBase { } } } + + private def withTinyPartitions(f: => Unit): Unit = { + withTempDir { dir => + val random = new Random(42) + val sizes = Seq.fill(700)(1 + random.nextInt(3)) ++ Seq(400) ++ Seq.fill(200)(1) + val positiveNaN = java.lang.Double.longBitsToDouble(0x7ff8000000000001L) + val negativeNaN = java.lang.Double.longBitsToDouble(0xfff8000000000002L) + val doubles = Seq(Some(0.0d), Some(-0.0d), Some(positiveNaN), Some(negativeNaN), None) + var id = 0 + val rows = sizes.zipWithIndex.flatMap { case (size, p) => + val k1 = if (p < 3) None else Some(p / 4) + val k2 = if (p % 4 == 0) None else Some(s"key-${p % 4}") + (0 until size).map { i => + id += 1 + val ts = + if (i == 0 && p % 5 == 0) None + else Some(new Timestamp(1700000000000L + i * 1000L)) + val eventTimeMs = if (random.nextInt(6) == 0) None else Some(random.nextInt(4).toLong) + val v = if (random.nextInt(5) == 0) None else Some(random.nextInt(100)) + val d = doubles((p / 7) % doubles.size) + (id, k1, k2, ts, eventTimeMs, v, s"payload-$id", d) + } + } + rows + .toDF("id", "k1", "k2", "effective_ts", "eventTimeMs", "v", "s", "d") + .repartition(3) + .write + .mode("overwrite") + .parquet(dir.toString) + spark.read.parquet(dir.toString).createOrReplaceTempView("tiny") + f + } + } + + private def checkAllModes(query: String): Unit = { + for { + sorted <- Seq("true", "false") + batchSize <- Seq("3", "8192") + } { + withSQLConf( + CometConf.COMET_EXEC_WINDOW_SORTED_ENABLED.key -> sorted, + CometConf.COMET_BATCH_SIZE.key -> batchSize) { + val (_, plan) = checkSparkAnswerAndOperator(sql(query)) + assert(collect(plan) { case w: CometWindowExec => w }.nonEmpty) + assert(collect(plan) { case w: SparkWindowExec => w }.isEmpty) + } + } + } + + test("sorted window: LEAD of a timestamp with a far-future default over tiny partitions") { + withTinyPartitions { + checkAllModes(""" + SELECT id, k1, k2, effective_ts, + lead(effective_ts, 1, TIMESTAMP'9999-12-31 23:59:59') + OVER (PARTITION BY k1, k2 ORDER BY effective_ts ASC NULLS FIRST + ROWS BETWEEN 1 FOLLOWING AND 1 FOLLOWING) AS next_ts + FROM tiny + """) + } + } + + test("sorted window: ROW_NUMBER DESC NULLS LAST, alone and filtered to the first row") { + withTinyPartitions { + checkAllModes(""" + SELECT id, k1, eventTimeMs, + row_number() OVER (PARTITION BY k1 ORDER BY eventTimeMs DESC NULLS LAST, id) AS rn + FROM tiny + """) + checkAllModes(""" + SELECT id, k1, k2, s FROM ( + SELECT *, + row_number() OVER (PARTITION BY k1, k2 ORDER BY eventTimeMs DESC NULLS LAST, id) AS rn + FROM tiny + ) WHERE rn = 1 + """) + } + } + + test("sorted window: ranking functions with ties and multi-column keys") { + withTinyPartitions { + checkAllModes(""" + SELECT id, k1, k2, eventTimeMs, + row_number() OVER w AS rn, + rank() OVER w AS r, + dense_rank() OVER w AS dr + FROM tiny + WINDOW w AS (PARTITION BY k1, k2 ORDER BY eventTimeMs DESC NULLS LAST, id) + """) + checkAllModes(""" + SELECT id, k1, eventTimeMs, + rank() OVER (PARTITION BY k1 ORDER BY eventTimeMs ASC NULLS FIRST) AS r, + dense_rank() OVER (PARTITION BY k1 ORDER BY eventTimeMs ASC NULLS FIRST) AS dr + FROM tiny + """) + checkAllModes(""" + SELECT id, k1, k2, eventTimeMs, + rank() OVER w AS r, + dense_rank() OVER w AS dr + FROM tiny + WINDOW w AS (PARTITION BY k1, k2 ORDER BY eventTimeMs DESC NULLS LAST) + """) + } + } + + test("sorted window: LEAD and LAG with offsets 1 to 3 and defaults") { + withTinyPartitions { + checkAllModes(""" + SELECT id, k1, k2, v, s, + lead(v) OVER w AS lead1, + lead(v, 2, -1) OVER w AS lead2, + lead(s, 3, 'none') OVER w AS lead3, + lag(v) OVER w AS lag1, + lag(v, 2, -2) OVER w AS lag2, + lag(s, 3) OVER w AS lag3, + row_number() OVER w AS rn + FROM tiny + WINDOW w AS (PARTITION BY k1, k2 ORDER BY id) + """) + } + } + + test("sorted window: struct, string and floating-point partition keys and nested values") { + withTinyPartitions { + checkAllModes(""" + SELECT id, d, v, + row_number() OVER (PARTITION BY d ORDER BY id) AS rn_d, + lead(named_struct('v', v, 's', array(s, k2))) + OVER (PARTITION BY d ORDER BY id) AS lead_nested, + row_number() OVER (PARTITION BY named_struct('a', k1 % 3, 'b', k2) ORDER BY id) AS rn_s + FROM tiny + """) + } + } + + test("sorted window: LEAD IGNORE NULLS and other functions keep the existing operators") { + withTinyPartitions { + checkAllModes(""" + SELECT id, k1, v, + lead(v) IGNORE NULLS OVER (PARTITION BY k1 ORDER BY id) AS lead_in, + row_number() OVER (PARTITION BY k1 ORDER BY id) AS rn + FROM tiny + """) + checkAllModes(""" + SELECT id, k1, v, + sum(v) OVER (PARTITION BY k1 ORDER BY id) AS running, + lag(v) OVER (PARTITION BY k1 ORDER BY id) AS previous + FROM tiny + """) + } + } }