diff --git a/datafusion/functions-aggregate/Cargo.toml b/datafusion/functions-aggregate/Cargo.toml index 5abea16e2cc81..04ba28daa722d 100644 --- a/datafusion/functions-aggregate/Cargo.toml +++ b/datafusion/functions-aggregate/Cargo.toml @@ -96,6 +96,10 @@ harness = false name = "percentile_cont" harness = false +[[bench]] +name = "any_value" +harness = false + [[bench]] name = "sliding_max" harness = false diff --git a/datafusion/functions-aggregate/benches/any_value.rs b/datafusion/functions-aggregate/benches/any_value.rs new file mode 100644 index 0000000000000..6f7096c1f9745 --- /dev/null +++ b/datafusion/functions-aggregate/benches/any_value.rs @@ -0,0 +1,130 @@ +// 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::hint::black_box; +use std::sync::Arc; + +use arrow::array::{ArrayRef, Int64Array, StringArray}; +use arrow::datatypes::{DataType, Field, Schema}; +use criterion::{BatchSize, Criterion, criterion_group, criterion_main}; +use datafusion_expr::function::AccumulatorArgs; +use datafusion_expr::{AggregateUDFImpl, EmitTo, GroupsAccumulator}; +use datafusion_functions_aggregate::any_value::AnyValue; +use datafusion_physical_expr::GroupsAccumulatorAdapter; +use datafusion_physical_expr::expressions::col; + +const BATCH_SIZE: usize = 8192; +const NUM_GROUPS: usize = 4096; + +fn with_accumulator_args( + data_type: DataType, + f: impl FnOnce(AccumulatorArgs<'_>) -> T, +) -> T { + let schema = Schema::new(vec![Field::new("value", data_type.clone(), true)]); + let expr = col("value", &schema).unwrap(); + let expr_fields = vec![expr.return_field(&schema).unwrap()]; + let exprs = vec![expr]; + let return_field = Field::new("any_value", data_type, true).into(); + + f(AccumulatorArgs { + return_field, + schema: &schema, + expr_fields: &expr_fields, + ignore_nulls: false, + order_bys: &[], + is_reversed: false, + name: "any_value(value)", + is_distinct: false, + exprs: &exprs, + }) +} + +fn native_accumulator(data_type: DataType) -> Box { + with_accumulator_args(data_type, |args| { + AnyValue::new().create_groups_accumulator(args).unwrap() + }) +} + +fn adapter_accumulator(data_type: DataType) -> Box { + Box::new(GroupsAccumulatorAdapter::new(move || { + with_accumulator_args(data_type.clone(), |args| AnyValue::new().accumulator(args)) + })) +} + +fn run_grouped( + accumulator: &mut dyn GroupsAccumulator, + values: &ArrayRef, + group_indices: &[usize], +) { + accumulator + .update_batch( + std::slice::from_ref(values), + group_indices, + None, + NUM_GROUPS, + ) + .unwrap(); + black_box(accumulator.evaluate(EmitTo::All).unwrap()); +} + +fn benchmark_type(c: &mut Criterion, name: &str, values: &ArrayRef) { + let group_indices = (0..BATCH_SIZE) + .map(|row| row % NUM_GROUPS) + .collect::>(); + let data_type = values.data_type().clone(); + + let mut group = c.benchmark_group(format!("any_value grouped {name}")); + group.bench_function("native", |b| { + b.iter_batched( + || native_accumulator(data_type.clone()), + |mut accumulator| { + run_grouped(accumulator.as_mut(), values, &group_indices); + }, + BatchSize::SmallInput, + ) + }); + group.bench_function("adapter", |b| { + b.iter_batched( + || adapter_accumulator(data_type.clone()), + |mut accumulator| { + run_grouped(accumulator.as_mut(), values, &group_indices); + }, + BatchSize::SmallInput, + ) + }); + group.finish(); +} + +fn criterion_benchmark(c: &mut Criterion) { + let int_values = Arc::new( + (0..BATCH_SIZE) + .map(|row| (row % 17 != 0).then_some(row as i64)) + .collect::(), + ) as ArrayRef; + benchmark_type(c, "int64", &int_values); + + let strings = (0..BATCH_SIZE) + .map(|row| (row % 17 != 0).then(|| format!("value-{row}"))) + .collect::>(); + let string_values = + Arc::new(StringArray::from_iter(strings.iter().map(Option::as_deref))) + as ArrayRef; + benchmark_type(c, "utf8", &string_values); +} + +criterion_group!(benches, criterion_benchmark); +criterion_main!(benches); diff --git a/datafusion/functions-aggregate/src/any_value.rs b/datafusion/functions-aggregate/src/any_value.rs index dc3bd23d806fc..5dc791e1d97bd 100644 --- a/datafusion/functions-aggregate/src/any_value.rs +++ b/datafusion/functions-aggregate/src/any_value.rs @@ -19,18 +19,33 @@ use std::fmt::Debug; use std::hash::Hash; +use std::mem::{size_of, size_of_val}; use std::sync::Arc; -use arrow::datatypes::{DataType, Field, FieldRef}; -use datafusion_common::{Result, not_impl_err}; +use arrow::array::{ + Array, ArrayRef, AsArray, BooleanArray, BooleanBufferBuilder, PrimitiveArray, +}; +use arrow::buffer::{BooleanBuffer, NullBuffer}; +use arrow::datatypes::{ + ArrowPrimitiveType, DataType, Date32Type, Date64Type, Decimal32Type, Decimal64Type, + Decimal128Type, Decimal256Type, Field, FieldRef, Float16Type, Float32Type, + Float64Type, Int8Type, Int16Type, Int32Type, Int64Type, Time32MillisecondType, + Time32SecondType, Time64MicrosecondType, Time64NanosecondType, TimeUnit, + TimestampMicrosecondType, TimestampMillisecondType, TimestampNanosecondType, + TimestampSecondType, UInt8Type, UInt16Type, UInt32Type, UInt64Type, +}; +use datafusion_common::cast::as_boolean_array; +use datafusion_common::{Result, ScalarValue, not_impl_err}; use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs}; use datafusion_expr::utils::{AggregateOrderSensitivity, format_state_name}; use datafusion_expr::{ - Accumulator, AggregateUDFImpl, Documentation, Signature, Volatility, + Accumulator, AggregateUDFImpl, Documentation, EmitTo, GroupsAccumulator, Signature, + Volatility, }; use datafusion_macros::user_doc; use crate::first_last::TrivialFirstValueAccumulator; +use crate::first_last::state::{BytesValueState, ValueState, take_need}; make_udaf_expr_and_func!( AnyValue, @@ -115,6 +130,17 @@ impl AggregateUDFImpl for AnyValue { ]) } + fn groups_accumulator_supported(&self, _args: AccumulatorArgs) -> bool { + true + } + + fn create_groups_accumulator( + &self, + args: AccumulatorArgs, + ) -> Result> { + create_groups_accumulator(args.return_field.data_type()) + } + fn order_sensitivity(&self) -> AggregateOrderSensitivity { AggregateOrderSensitivity::Insensitive } @@ -123,3 +149,566 @@ impl AggregateUDFImpl for AnyValue { self.doc() } } +fn create_groups_accumulator(data_type: &DataType) -> Result> { + macro_rules! instantiate_primitive { + ($t:ty) => { + Ok(Box::new(PrimitiveAnyValueGroupsAccumulator::<$t>::new( + data_type.clone(), + )) as _) + }; + } + + match data_type { + DataType::Int8 => instantiate_primitive!(Int8Type), + DataType::Int16 => instantiate_primitive!(Int16Type), + DataType::Int32 => instantiate_primitive!(Int32Type), + DataType::Int64 => instantiate_primitive!(Int64Type), + DataType::UInt8 => instantiate_primitive!(UInt8Type), + DataType::UInt16 => instantiate_primitive!(UInt16Type), + DataType::UInt32 => instantiate_primitive!(UInt32Type), + DataType::UInt64 => instantiate_primitive!(UInt64Type), + DataType::Float16 => instantiate_primitive!(Float16Type), + DataType::Float32 => instantiate_primitive!(Float32Type), + DataType::Float64 => instantiate_primitive!(Float64Type), + + DataType::Decimal32(_, _) => instantiate_primitive!(Decimal32Type), + DataType::Decimal64(_, _) => instantiate_primitive!(Decimal64Type), + DataType::Decimal128(_, _) => instantiate_primitive!(Decimal128Type), + DataType::Decimal256(_, _) => instantiate_primitive!(Decimal256Type), + + DataType::Timestamp(TimeUnit::Second, _) => { + instantiate_primitive!(TimestampSecondType) + } + DataType::Timestamp(TimeUnit::Millisecond, _) => { + instantiate_primitive!(TimestampMillisecondType) + } + DataType::Timestamp(TimeUnit::Microsecond, _) => { + instantiate_primitive!(TimestampMicrosecondType) + } + DataType::Timestamp(TimeUnit::Nanosecond, _) => { + instantiate_primitive!(TimestampNanosecondType) + } + + DataType::Date32 => instantiate_primitive!(Date32Type), + DataType::Date64 => instantiate_primitive!(Date64Type), + DataType::Time32(TimeUnit::Second) => instantiate_primitive!(Time32SecondType), + DataType::Time32(TimeUnit::Millisecond) => { + instantiate_primitive!(Time32MillisecondType) + } + DataType::Time64(TimeUnit::Microsecond) => { + instantiate_primitive!(Time64MicrosecondType) + } + DataType::Time64(TimeUnit::Nanosecond) => { + instantiate_primitive!(Time64NanosecondType) + } + + DataType::Utf8 + | DataType::LargeUtf8 + | DataType::Utf8View + | DataType::Binary + | DataType::LargeBinary + | DataType::BinaryView => Ok(Box::new(BytesAnyValueGroupsAccumulator::try_new( + data_type.clone(), + )?) as _), + + _ => Ok(Box::new(AnyValueGroupsAccumulator::try_new(data_type)?) as _), + } +} + +#[derive(Debug)] +struct BytesAnyValueGroupsAccumulator { + values: BytesValueState, + is_set: BooleanBufferBuilder, +} + +impl BytesAnyValueGroupsAccumulator { + fn try_new(data_type: DataType) -> Result { + Ok(Self { + values: BytesValueState::try_new(data_type)?, + is_set: BooleanBufferBuilder::new(0), + }) + } + + fn ensure_groups(&mut self, total_num_groups: usize) { + if self.is_set.len() < total_num_groups { + self.values.resize(total_num_groups); + self.is_set.resize(total_num_groups); + } + } + + fn take_state(&mut self, emit_to: EmitTo) -> Result<(ArrayRef, BooleanBuffer)> { + let values = self.values.take(emit_to)?; + let is_set = take_need(&mut self.is_set, emit_to); + Ok((values, is_set)) + } +} + +impl GroupsAccumulator for BytesAnyValueGroupsAccumulator { + fn update_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if opt_filter.is_none_or(|filter| filter.is_valid(row) && filter.value(row)) + && !self.is_set.get_bit(group_index) + && values[0].is_valid(row) + { + self.values.update(group_index, &values[0], row)?; + self.is_set.set_bit(group_index, true); + } + } + Ok(()) + } + + fn evaluate(&mut self, emit_to: EmitTo) -> Result { + self.take_state(emit_to).map(|(values, _)| values) + } + + fn state(&mut self, emit_to: EmitTo) -> Result> { + let (values, is_set) = self.take_state(emit_to)?; + Ok(vec![values, Arc::new(BooleanArray::new(is_set, None))]) + } + + fn merge_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 2, "any_value expects value and is_set state"); + let is_set = as_boolean_array(&values[1])?; + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if is_set.is_valid(row) + && is_set.value(row) + && !self.is_set.get_bit(group_index) + && values[0].is_valid(row) + { + self.values.update(group_index, &values[0], row)?; + self.is_set.set_bit(group_index, true); + } + } + Ok(()) + } + + fn convert_to_state( + &self, + values: &[ArrayRef], + opt_filter: Option<&BooleanArray>, + ) -> Result> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + let values = &values[0]; + let is_set = BooleanArray::from_iter((0..values.len()).map(|row| { + values.is_valid(row) + && opt_filter + .is_none_or(|filter| filter.is_valid(row) && filter.value(row)) + })); + Ok(vec![Arc::clone(values), Arc::new(is_set)]) + } + + fn size(&self) -> usize { + size_of_val(self) + self.values.size() + self.is_set.capacity() / 8 + } +} + +#[derive(Debug)] +struct PrimitiveAnyValueGroupsAccumulator { + values: Vec, + is_set: BooleanBufferBuilder, + data_type: DataType, +} + +impl PrimitiveAnyValueGroupsAccumulator { + fn new(data_type: DataType) -> Self { + Self { + values: vec![], + is_set: BooleanBufferBuilder::new(0), + data_type, + } + } + + fn ensure_groups(&mut self, total_num_groups: usize) { + if self.values.len() < total_num_groups { + self.values.resize(total_num_groups, T::default_value()); + self.is_set.resize(total_num_groups); + } + } + + fn take_state(&mut self, emit_to: EmitTo) -> (Vec, BooleanBuffer) { + let values = emit_to.take_needed(&mut self.values); + let is_set = self.is_set.finish(); + match emit_to { + EmitTo::All => (values, is_set), + EmitTo::First(n) => { + let emitted = is_set.slice(0, n); + self.is_set + .append_buffer(&is_set.slice(n, is_set.len() - n)); + (values, emitted) + } + } + } + + fn values_array( + &self, + values: Vec, + is_set: BooleanBuffer, + ) -> PrimitiveArray { + PrimitiveArray::::new(values.into(), Some(NullBuffer::new(is_set))) + .with_data_type(self.data_type.clone()) + } +} + +impl GroupsAccumulator + for PrimitiveAnyValueGroupsAccumulator +{ + fn update_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + let values = values[0].as_primitive::(); + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if opt_filter.is_none_or(|filter| filter.is_valid(row) && filter.value(row)) + && !self.is_set.get_bit(group_index) + && values.is_valid(row) + { + self.values[group_index] = values.value(row); + self.is_set.set_bit(group_index, true); + } + } + Ok(()) + } + + fn evaluate(&mut self, emit_to: EmitTo) -> Result { + let (values, is_set) = self.take_state(emit_to); + Ok(Arc::new(self.values_array(values, is_set))) + } + + fn state(&mut self, emit_to: EmitTo) -> Result> { + let (values, is_set) = self.take_state(emit_to); + Ok(vec![ + Arc::new(self.values_array(values, is_set.clone())), + Arc::new(BooleanArray::new(is_set, None)), + ]) + } + + fn merge_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 2, "any_value expects value and is_set state"); + let state_values = values[0].as_primitive::(); + let is_set = as_boolean_array(&values[1])?; + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if is_set.is_valid(row) + && is_set.value(row) + && !self.is_set.get_bit(group_index) + && state_values.is_valid(row) + { + self.values[group_index] = state_values.value(row); + self.is_set.set_bit(group_index, true); + } + } + Ok(()) + } + + fn convert_to_state( + &self, + values: &[ArrayRef], + opt_filter: Option<&BooleanArray>, + ) -> Result> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + let values = &values[0]; + let is_set = BooleanArray::from_iter((0..values.len()).map(|row| { + values.is_valid(row) + && opt_filter + .is_none_or(|filter| filter.is_valid(row) && filter.value(row)) + })); + Ok(vec![Arc::clone(values), Arc::new(is_set)]) + } + + fn size(&self) -> usize { + size_of_val(self) + + self.values.capacity() * size_of::() + + self.is_set.capacity() / 8 + } +} + +#[derive(Debug)] +struct AnyValueGroupsAccumulator { + values: Vec, + is_set: BooleanBufferBuilder, + null_value: ScalarValue, +} + +impl AnyValueGroupsAccumulator { + fn try_new(data_type: &DataType) -> Result { + Ok(Self { + values: vec![], + is_set: BooleanBufferBuilder::new(0), + null_value: ScalarValue::try_from(data_type)?, + }) + } + + fn ensure_groups(&mut self, total_num_groups: usize) { + if self.values.len() < total_num_groups { + self.values + .resize(total_num_groups, self.null_value.clone()); + self.is_set.resize(total_num_groups); + } + } + + fn take_state(&mut self, emit_to: EmitTo) -> (Vec, BooleanBuffer) { + let values = emit_to.take_needed(&mut self.values); + let is_set = self.is_set.finish(); + match emit_to { + EmitTo::All => (values, is_set), + EmitTo::First(n) => { + let emitted = is_set.slice(0, n); + self.is_set + .append_buffer(&is_set.slice(n, is_set.len() - n)); + (values, emitted) + } + } + } + + fn update_row( + &mut self, + values: &ArrayRef, + group_index: usize, + row: usize, + ) -> Result<()> { + if !self.is_set.get_bit(group_index) && values.is_valid(row) { + self.values[group_index] = ScalarValue::try_from_array(values, row)?; + self.is_set.set_bit(group_index, true); + } + Ok(()) + } +} + +impl GroupsAccumulator for AnyValueGroupsAccumulator { + fn update_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + let values = &values[0]; + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if opt_filter.is_none_or(|filter| filter.is_valid(row) && filter.value(row)) { + self.update_row(values, group_index, row)?; + } + } + Ok(()) + } + + fn evaluate(&mut self, emit_to: EmitTo) -> Result { + let (values, _) = self.take_state(emit_to); + ScalarValue::iter_to_array(values) + } + + fn state(&mut self, emit_to: EmitTo) -> Result> { + let (values, is_set) = self.take_state(emit_to); + Ok(vec![ + ScalarValue::iter_to_array(values)?, + Arc::new(BooleanArray::new(is_set, None)), + ]) + } + + fn merge_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + total_num_groups: usize, + ) -> Result<()> { + assert_eq!(values.len(), 2, "any_value expects value and is_set state"); + let is_set = as_boolean_array(&values[1])?; + self.ensure_groups(total_num_groups); + + for (row, &group_index) in group_indices.iter().enumerate() { + if is_set.is_valid(row) + && is_set.value(row) + && !self.is_set.get_bit(group_index) + { + self.values[group_index] = ScalarValue::try_from_array(&values[0], row)?; + self.is_set.set_bit(group_index, true); + } + } + Ok(()) + } + + fn convert_to_state( + &self, + values: &[ArrayRef], + opt_filter: Option<&BooleanArray>, + ) -> Result> { + assert_eq!(values.len(), 1, "any_value expects one argument"); + let values = &values[0]; + let is_set = BooleanArray::from_iter((0..values.len()).map(|row| { + values.is_valid(row) + && opt_filter + .is_none_or(|filter| filter.is_valid(row) && filter.value(row)) + })); + Ok(vec![Arc::clone(values), Arc::new(is_set)]) + } + + fn size(&self) -> usize { + size_of_val(self) + + ScalarValue::size_of_vec(&self.values) + + self.is_set.capacity() / 8 + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{BinaryViewArray, Int64Array, StringArray, StringViewArray}; + + #[test] + fn groups_accumulator_uses_first_non_null_value() -> Result<()> { + let mut acc = create_groups_accumulator(&DataType::Int64)?; + let values = Arc::new(Int64Array::from(vec![ + None, + Some(10), + Some(11), + Some(20), + None, + Some(30), + ])) as ArrayRef; + let filter = BooleanArray::from(vec![true, true, true, false, true, true]); + + acc.update_batch(&[values], &[0, 0, 0, 1, 1, 2], Some(&filter), 4)?; + let result = acc.evaluate(EmitTo::All)?; + let expected = + Arc::new(Int64Array::from(vec![Some(10), None, Some(30), None])) as ArrayRef; + assert_eq!(&result, &expected); + Ok(()) + } + + #[test] + fn groups_accumulator_merges_partial_state() -> Result<()> { + let mut acc = AnyValueGroupsAccumulator::try_new(&DataType::Utf8)?; + let values = Arc::new(StringArray::from(vec![Some("a"), Some("b"), Some("c")])) + as ArrayRef; + let is_set = Arc::new(BooleanArray::from(vec![false, true, true])) as ArrayRef; + + acc.merge_batch(&[values, is_set], &[0, 0, 1], 3)?; + let state = acc.state(EmitTo::All)?; + let expected_values = + Arc::new(StringArray::from(vec![Some("b"), Some("c"), None])) as ArrayRef; + let expected_is_set = + Arc::new(BooleanArray::from(vec![true, true, false])) as ArrayRef; + assert_eq!(&state[0], &expected_values); + assert_eq!(&state[1], &expected_is_set); + Ok(()) + } + + #[test] + fn groups_accumulator_convert_to_state_applies_filter_and_nulls() -> Result<()> { + let acc = create_groups_accumulator(&DataType::Int64)?; + let values = Arc::new(Int64Array::from(vec![Some(1), None, Some(3)])) as ArrayRef; + let filter = BooleanArray::from(vec![true, true, false]); + + let state = acc.convert_to_state(&[Arc::clone(&values)], Some(&filter))?; + let expected_is_set = + Arc::new(BooleanArray::from(vec![true, false, false])) as ArrayRef; + assert_eq!(&state[0], &values); + assert_eq!(&state[1], &expected_is_set); + Ok(()) + } + + #[test] + fn groups_accumulator_emit_first_retains_remaining_groups() -> Result<()> { + let mut acc = create_groups_accumulator(&DataType::Int64)?; + let values = + Arc::new(Int64Array::from(vec![Some(10), Some(20), Some(30)])) as ArrayRef; + acc.update_batch(&[values], &[0, 1, 2], None, 3)?; + + let first = acc.evaluate(EmitTo::First(2))?; + let expected_first = + Arc::new(Int64Array::from(vec![Some(10), Some(20)])) as ArrayRef; + assert_eq!(&first, &expected_first); + + let values = Arc::new(Int64Array::from(vec![Some(31), Some(40)])) as ArrayRef; + acc.update_batch(&[values], &[0, 1], None, 2)?; + let remaining = acc.evaluate(EmitTo::All)?; + let expected_remaining = + Arc::new(Int64Array::from(vec![Some(30), Some(40)])) as ArrayRef; + assert_eq!(&remaining, &expected_remaining); + Ok(()) + } + + #[test] + fn primitive_groups_accumulator_merges_partial_state() -> Result<()> { + let mut partial = create_groups_accumulator(&DataType::Int64)?; + let values = Arc::new(Int64Array::from(vec![Some(10), None, Some(30), Some(40)])) + as ArrayRef; + partial.update_batch(&[values], &[0, 1, 2, 2], None, 3)?; + let state = partial.state(EmitTo::All)?; + + let mut merged = create_groups_accumulator(&DataType::Int64)?; + merged.merge_batch(&state, &[0, 1, 2], 4)?; + let result = merged.evaluate(EmitTo::All)?; + let expected = + Arc::new(Int64Array::from(vec![Some(10), None, Some(30), None])) as ArrayRef; + assert_eq!(&result, &expected); + Ok(()) + } + + #[test] + fn bytes_groups_accumulator_supports_view_types_and_emit_first() -> Result<()> { + let mut strings = create_groups_accumulator(&DataType::Utf8View)?; + let values = Arc::new(StringViewArray::from(vec![ + Some("first"), + Some("ignored"), + Some("remaining"), + ])) as ArrayRef; + strings.update_batch(&[values], &[0, 0, 1], None, 2)?; + + let first = strings.evaluate(EmitTo::First(1))?; + let expected_first = + Arc::new(StringViewArray::from(vec![Some("first")])) as ArrayRef; + assert_eq!(&first, &expected_first); + + let remaining = strings.evaluate(EmitTo::All)?; + let expected_remaining = + Arc::new(StringViewArray::from(vec![Some("remaining")])) as ArrayRef; + assert_eq!(&remaining, &expected_remaining); + + let mut binary = create_groups_accumulator(&DataType::BinaryView)?; + let values = Arc::new(BinaryViewArray::from(vec![ + Some(b"a".as_slice()), + None, + Some(b"b".as_slice()), + ])) as ArrayRef; + binary.update_batch(&[values], &[0, 1, 1], None, 3)?; + let result = binary.evaluate(EmitTo::All)?; + let expected = Arc::new(BinaryViewArray::from(vec![ + Some(b"a".as_slice()), + Some(b"b".as_slice()), + None, + ])) as ArrayRef; + assert_eq!(&result, &expected); + Ok(()) + } +} diff --git a/datafusion/functions-aggregate/src/first_last.rs b/datafusion/functions-aggregate/src/first_last.rs index ea45e42e84f33..f8ee8355289cb 100644 --- a/datafusion/functions-aggregate/src/first_last.rs +++ b/datafusion/functions-aggregate/src/first_last.rs @@ -49,7 +49,7 @@ use datafusion_functions_aggregate_common::utils::get_sort_options; use datafusion_macros::user_doc; use datafusion_physical_expr_common::sort_expr::LexOrdering; -mod state; +pub(crate) mod state; use state::{BytesValueState, PrimitiveValueState, ValueState}; diff --git a/datafusion/functions-aggregate/src/first_last/state.rs b/datafusion/functions-aggregate/src/first_last/state.rs index cd7114bf04f9c..c45cfd8fc841f 100644 --- a/datafusion/functions-aggregate/src/first_last/state.rs +++ b/datafusion/functions-aggregate/src/first_last/state.rs @@ -111,6 +111,7 @@ impl ValueState for PrimitiveValueState { /// in the input, while `BytesValueState` needs to support setting `NULL` values /// to correctly implement `RESPECT NULLS` behavior. /// +#[derive(Debug)] pub(crate) struct BytesValueState { vals: Vec>>, data_type: DataType,