Skip to main content

datafusion_functions_aggregate/
array_agg.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! `ARRAY_AGG` aggregate implementation: [`ArrayAgg`]
19
20use std::cmp::Ordering;
21use std::collections::VecDeque;
22use std::mem::{size_of, size_of_val, take};
23use std::sync::Arc;
24
25use arrow::array::{
26    Array, ArrayRef, AsArray, BooleanArray, ListArray, NullBufferBuilder, StructArray,
27    UInt32Array, new_empty_array,
28};
29use arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
30use arrow::compute::{SortOptions, cast, filter};
31use arrow::datatypes::{DataType, Field, FieldRef, Fields};
32use arrow::row::{OwnedRow, Row, RowConverter, Rows, SortField};
33
34use datafusion_common::cast::as_list_array;
35use datafusion_common::hash_utils::{RandomState, create_hashes};
36use datafusion_common::utils::proxy::HashTableAllocExt;
37use datafusion_common::utils::{
38    SingleRowListArrayBuilder, compare_rows, get_row_at_idx, take_function_args,
39};
40use datafusion_common::{
41    Result, ScalarValue, assert_eq_or_internal_err, exec_err, internal_err,
42};
43use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs};
44use datafusion_expr::utils::format_state_name;
45use datafusion_expr::{
46    Accumulator, AggregateUDFImpl, Documentation, EmitTo, GroupsAccumulator, Signature,
47    Volatility,
48};
49use datafusion_functions_aggregate_common::aggregate::groups_accumulator::nulls::filter_to_nulls;
50use datafusion_functions_aggregate_common::merge_arrays::merge_ordered_arrays;
51use datafusion_functions_aggregate_common::order::AggregateOrderSensitivity;
52use datafusion_functions_aggregate_common::utils::ordering_fields;
53use datafusion_macros::user_doc;
54use datafusion_physical_expr_common::sort_expr::{LexOrdering, PhysicalSortExpr};
55use hashbrown::hash_table::HashTable;
56
57make_udaf_expr_and_func!(
58    ArrayAgg,
59    array_agg,
60    expression,
61    "input values, including nulls, concatenated into an array",
62    array_agg_udaf
63);
64
65#[user_doc(
66    doc_section(label = "General Functions"),
67    description = r#"Returns an array created from the expression elements. If ordering is required, elements are inserted in the specified order.
68This aggregation function can only mix DISTINCT and ORDER BY if the ordering expression is exactly the same as the argument expression."#,
69    syntax_example = "array_agg(expression [ORDER BY expression])",
70    sql_example = r#"
71```sql
72> SELECT array_agg(column_name ORDER BY other_column) FROM table_name;
73+-----------------------------------------------+
74| array_agg(column_name ORDER BY other_column)  |
75+-----------------------------------------------+
76| [element1, element2, element3]                |
77+-----------------------------------------------+
78> SELECT array_agg(DISTINCT column_name ORDER BY column_name) FROM table_name;
79+--------------------------------------------------------+
80| array_agg(DISTINCT column_name ORDER BY column_name)  |
81+--------------------------------------------------------+
82| [element1, element2, element3]                         |
83+--------------------------------------------------------+
84```
85"#,
86    standard_argument(name = "expression",)
87)]
88#[derive(Debug, PartialEq, Eq, Hash)]
89/// ARRAY_AGG aggregate expression
90pub struct ArrayAgg {
91    signature: Signature,
92    is_input_pre_ordered: bool,
93}
94
95impl Default for ArrayAgg {
96    fn default() -> Self {
97        Self {
98            signature: Signature::any(1, Volatility::Immutable),
99            is_input_pre_ordered: false,
100        }
101    }
102}
103
104impl AggregateUDFImpl for ArrayAgg {
105    fn name(&self) -> &str {
106        "array_agg"
107    }
108
109    fn signature(&self) -> &Signature {
110        &self.signature
111    }
112
113    fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
114        Ok(DataType::List(Arc::new(Field::new_list_field(
115            arg_types[0].clone(),
116            true,
117        ))))
118    }
119
120    fn state_fields(&self, args: StateFieldsArgs) -> Result<Vec<FieldRef>> {
121        if args.is_distinct {
122            return Ok(vec![
123                Field::new_list(
124                    format_state_name(args.name, "distinct_array_agg"),
125                    // See COMMENTS.md to understand why nullable is set to true
126                    Field::new_list_field(args.input_fields[0].data_type().clone(), true),
127                    true,
128                )
129                .into(),
130            ]);
131        }
132
133        let mut fields = vec![
134            Field::new_list(
135                format_state_name(args.name, "array_agg"),
136                // See COMMENTS.md to understand why nullable is set to true
137                Field::new_list_field(args.input_fields[0].data_type().clone(), true),
138                true,
139            )
140            .into(),
141        ];
142
143        if args.ordering_fields.is_empty() {
144            return Ok(fields);
145        }
146
147        let orderings = args.ordering_fields.to_vec();
148        fields.push(
149            Field::new_list(
150                format_state_name(args.name, "array_agg_orderings"),
151                Field::new_list_field(DataType::Struct(Fields::from(orderings)), true),
152                false,
153            )
154            .into(),
155        );
156
157        Ok(fields)
158    }
159
160    fn order_sensitivity(&self) -> AggregateOrderSensitivity {
161        AggregateOrderSensitivity::SoftRequirement
162    }
163
164    fn with_beneficial_ordering(
165        self: Arc<Self>,
166        beneficial_ordering: bool,
167    ) -> Result<Option<Arc<dyn AggregateUDFImpl>>> {
168        Ok(Some(Arc::new(Self {
169            signature: self.signature.clone(),
170            is_input_pre_ordered: beneficial_ordering,
171        })))
172    }
173
174    fn accumulator(&self, acc_args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
175        let field = &acc_args.expr_fields[0];
176        let data_type = field.data_type();
177        let ignore_nulls = acc_args.ignore_nulls && field.is_nullable();
178
179        if acc_args.is_distinct {
180            // Limitation similar to Postgres. The aggregation function can only mix
181            // DISTINCT and ORDER BY if all the expressions in the ORDER BY appear
182            // also in the arguments of the function. This implies that if the
183            // aggregation function only accepts one argument, only one argument
184            // can be used in the ORDER BY, For example:
185            //
186            // ARRAY_AGG(DISTINCT col)
187            //
188            // can only be mixed with an ORDER BY if the order expression is "col".
189            //
190            // ARRAY_AGG(DISTINCT col ORDER BY col)                         <- Valid
191            // ARRAY_AGG(DISTINCT concat(col, '') ORDER BY concat(col, '')) <- Valid
192            // ARRAY_AGG(DISTINCT col ORDER BY other_col)                   <- Invalid
193            // ARRAY_AGG(DISTINCT col ORDER BY concat(col, ''))             <- Invalid
194            let sort_option = match acc_args.order_bys {
195                [single] if single.expr.eq(&acc_args.exprs[0]) => Some(single.options),
196                [] => None,
197                _ => {
198                    return exec_err!(
199                        "In an aggregate with DISTINCT, ORDER BY expressions must appear in argument list"
200                    );
201                }
202            };
203            return Ok(Box::new(DistinctArrayAggAccumulator::try_new(
204                data_type,
205                sort_option,
206                ignore_nulls,
207            )?));
208        }
209
210        let Some(ordering) = LexOrdering::new(acc_args.order_bys.to_vec()) else {
211            return Ok(Box::new(ArrayAggAccumulator::try_new(
212                data_type,
213                ignore_nulls,
214            )?));
215        };
216
217        let ordering_dtypes = ordering
218            .iter()
219            .map(|e| e.expr.data_type(acc_args.schema))
220            .collect::<Result<Vec<_>>>()?;
221
222        OrderSensitiveArrayAggAccumulator::try_new(
223            data_type,
224            &ordering_dtypes,
225            ordering,
226            self.is_input_pre_ordered,
227            acc_args.is_reversed,
228            ignore_nulls,
229        )
230        .map(|acc| Box::new(acc) as _)
231    }
232
233    fn reverse_expr(&self) -> datafusion_expr::ReversedUDAF {
234        datafusion_expr::ReversedUDAF::Reversed(array_agg_udaf())
235    }
236
237    fn groups_accumulator_supported(&self, args: AccumulatorArgs) -> bool {
238        !args.is_distinct && args.order_bys.is_empty()
239    }
240
241    fn create_groups_accumulator(
242        &self,
243        args: AccumulatorArgs,
244    ) -> Result<Box<dyn GroupsAccumulator>> {
245        let field = &args.expr_fields[0];
246        let data_type = field.data_type().clone();
247        let ignore_nulls = args.ignore_nulls && field.is_nullable();
248        Ok(Box::new(ArrayAggGroupsAccumulator::new(
249            data_type,
250            ignore_nulls,
251        )))
252    }
253
254    fn supports_null_handling_clause(&self) -> bool {
255        true
256    }
257
258    fn documentation(&self) -> Option<&Documentation> {
259        self.doc()
260    }
261}
262
263#[derive(Debug)]
264pub struct ArrayAggAccumulator {
265    values: VecDeque<ArrayRef>,
266    datatype: DataType,
267    ignore_nulls: bool,
268    /// Number of elements already consumed (retracted) from the front array.
269    /// Used by sliding window frames to avoid copying on partial retract.
270    front_offset: usize,
271}
272
273impl ArrayAggAccumulator {
274    /// new array_agg accumulator based on given item data type
275    pub fn try_new(datatype: &DataType, ignore_nulls: bool) -> Result<Self> {
276        Ok(Self {
277            values: VecDeque::new(),
278            datatype: datatype.clone(),
279            ignore_nulls,
280            front_offset: 0,
281        })
282    }
283
284    /// This function will return the underlying list array values if all valid values are consecutive without gaps (i.e. no null value point to a non-empty list)
285    /// If there are gaps but only in the end of the list array, the function will return the values without the null values in the end
286    fn get_optional_values_to_merge_as_is(list_array: &ListArray) -> Option<ArrayRef> {
287        let offsets = list_array.value_offsets();
288        // Offsets always have at least 1 value
289        let initial_offset = offsets[0];
290        let null_count = list_array.null_count();
291
292        // If no nulls than just use the fast path
293        // This is ok as the state is a ListArray rather than a ListViewArray so all the values are consecutive
294        if null_count == 0 {
295            // According to Arrow specification, the first offset can be non-zero
296            let list_values = list_array.values().slice(
297                initial_offset as usize,
298                (offsets[offsets.len() - 1] - initial_offset) as usize,
299            );
300            return Some(list_values);
301        }
302
303        // If all the values are null than just return an empty values array
304        if list_array.null_count() == list_array.len() {
305            return Some(list_array.values().slice(0, 0));
306        }
307
308        // According to the Arrow spec, null values can point to non-empty lists
309        // So this will check if all null values starting from the first valid value to the last one point to a 0 length list so we can just slice the underlying value
310
311        // Unwrapping is safe as we just checked if there is a null value
312        let nulls = list_array.nulls().unwrap();
313
314        let mut valid_slices_iter = nulls.valid_slices();
315
316        // This is safe as we validated that there is at least 1 valid value in the array
317        let (start, end) = valid_slices_iter.next().unwrap();
318
319        let start_offset = offsets[start];
320
321        // End is exclusive, so it already point to the last offset value
322        // This is valid as the length of the array is always 1 less than the length of the offsets
323        let mut end_offset_of_last_valid_value = offsets[end];
324
325        for (start, end) in valid_slices_iter {
326            // If there is a null value that point to a non-empty list than the start offset of the valid value
327            // will be different that the end offset of the last valid value
328            if offsets[start] != end_offset_of_last_valid_value {
329                return None;
330            }
331
332            // End is exclusive, so it already point to the last offset value
333            // This is valid as the length of the array is always 1 less than the length of the offsets
334            end_offset_of_last_valid_value = offsets[end];
335        }
336
337        let consecutive_valid_values = list_array.values().slice(
338            start_offset as usize,
339            (end_offset_of_last_valid_value - start_offset) as usize,
340        );
341
342        Some(consecutive_valid_values)
343    }
344}
345
346impl Accumulator for ArrayAggAccumulator {
347    fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
348        // Append value like Int64Array(1,2,3)
349        if values.is_empty() {
350            return Ok(());
351        }
352
353        assert_eq_or_internal_err!(values.len(), 1, "expects single batch");
354
355        let val = &values[0];
356        let nulls = if self.ignore_nulls {
357            val.logical_nulls()
358        } else {
359            None
360        };
361
362        let val = match nulls {
363            Some(nulls) if nulls.null_count() >= val.len() => return Ok(()),
364            Some(nulls) => filter(val, &BooleanArray::new(nulls.inner().clone(), None))?,
365            None => Arc::clone(val),
366        };
367
368        if !val.is_empty() {
369            self.values.push_back(val)
370        }
371
372        Ok(())
373    }
374
375    fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
376        // Append value like ListArray(Int64Array(1,2,3), Int64Array(4,5,6))
377        if states.is_empty() {
378            return Ok(());
379        }
380
381        assert_eq_or_internal_err!(states.len(), 1, "expects single state");
382
383        let list_arr = as_list_array(&states[0])?;
384
385        match Self::get_optional_values_to_merge_as_is(list_arr) {
386            Some(values) => {
387                // Make sure we don't insert empty lists
388                if !values.is_empty() {
389                    self.values.push_back(values);
390                }
391            }
392            None => {
393                for arr in list_arr.iter().flatten() {
394                    self.values.push_back(arr);
395                }
396            }
397        }
398
399        Ok(())
400    }
401
402    fn state(&mut self) -> Result<Vec<ScalarValue>> {
403        Ok(vec![self.evaluate()?])
404    }
405
406    fn evaluate(&mut self) -> Result<ScalarValue> {
407        if self.values.is_empty() {
408            return Ok(ScalarValue::new_null_list(self.datatype.clone(), true, 1));
409        }
410
411        let element_arrays: Vec<ArrayRef> = self
412            .values
413            .iter()
414            .enumerate()
415            .map(|(i, a)| {
416                if i == 0 && self.front_offset > 0 {
417                    a.slice(self.front_offset, a.len() - self.front_offset)
418                } else {
419                    Arc::clone(a)
420                }
421            })
422            .collect();
423
424        let element_refs: Vec<&dyn Array> =
425            element_arrays.iter().map(|a| a.as_ref()).collect();
426
427        if element_refs.iter().all(|a| a.is_empty()) {
428            return Ok(ScalarValue::new_null_list(self.datatype.clone(), true, 1));
429        }
430
431        let concated_array = arrow::compute::concat(&element_refs)?;
432
433        Ok(SingleRowListArrayBuilder::new(concated_array).build_list_scalar())
434    }
435
436    fn retract_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
437        if values.is_empty() {
438            return Ok(());
439        }
440
441        assert_eq_or_internal_err!(values.len(), 1, "expects single batch");
442
443        let val = &values[0];
444        let mut to_retract = if self.ignore_nulls {
445            val.len() - val.logical_null_count()
446        } else {
447            val.len()
448        };
449
450        while to_retract > 0 {
451            let Some(front) = self.values.front() else {
452                break;
453            };
454            let available = front.len() - self.front_offset;
455            if to_retract >= available {
456                self.values.pop_front();
457                to_retract -= available;
458                self.front_offset = 0;
459            } else {
460                self.front_offset += to_retract;
461                to_retract = 0;
462            }
463        }
464
465        Ok(())
466    }
467
468    fn supports_retract_batch(&self) -> bool {
469        true
470    }
471
472    fn size(&self) -> usize {
473        size_of_val(self)
474            + (size_of::<ArrayRef>() * self.values.capacity())
475            + self
476                .values
477                .iter()
478                // Each ArrayRef might be just a reference to a bigger array, and many
479                // ArrayRefs here might be referencing exactly the same array, so if we
480                // were to call `arr.get_array_memory_size()`, we would be double-counting
481                // the same underlying data many times.
482                //
483                // Instead, we do an approximation by estimating how much memory each
484                // ArrayRef would occupy if its underlying data was fully owned by this
485                // accumulator.
486                //
487                // Note that this is just an estimation, but the reality is that this
488                // accumulator might not own any data.
489                .map(|arr| arr.to_data().get_slice_memory_size().unwrap_or_default())
490                .sum::<usize>()
491            + self.datatype.size()
492            - size_of_val(&self.datatype)
493    }
494}
495
496#[derive(Debug)]
497struct ArrayAggGroupsAccumulator {
498    datatype: DataType,
499    ignore_nulls: bool,
500    /// Source arrays — input arrays (from update_batch) or list backing
501    /// arrays (from merge_batch).
502    batches: Vec<ArrayRef>,
503    /// Per-batch list of (group_idx, row_idx) pairs.
504    batch_entries: Vec<Vec<(u32, u32)>>,
505    /// Total number of groups tracked.
506    num_groups: usize,
507}
508
509impl ArrayAggGroupsAccumulator {
510    fn new(datatype: DataType, ignore_nulls: bool) -> Self {
511        Self {
512            datatype,
513            ignore_nulls,
514            batches: Vec::new(),
515            batch_entries: Vec::new(),
516            num_groups: 0,
517        }
518    }
519
520    fn clear_state(&mut self) {
521        // `size()` measures Vec capacity rather than len, so allocate new
522        // buffers instead of using `clear()`.
523        self.batches = Vec::new();
524        self.batch_entries = Vec::new();
525        self.num_groups = 0;
526    }
527
528    fn compact_retained_state(&mut self, emit_groups: usize) -> Result<()> {
529        // EmitTo::First is used to recover from memory pressure. Simply
530        // removing emitted entries in place is not enough because mixed batches
531        // would continue to pin their original Array arrays, even if only a few
532        // retained rows remain.
533        //
534        // Rebuild the retained state from scratch so fully emitted batches are
535        // dropped, mixed batches are compacted to arrays containing only the
536        // surviving rows, and retained metadata is right-sized.
537        let emit_groups = emit_groups as u32;
538        let old_batches = take(&mut self.batches);
539        let old_batch_entries = take(&mut self.batch_entries);
540
541        let mut batches = Vec::new();
542        let mut batch_entries = Vec::new();
543
544        for (batch, entries) in old_batches.into_iter().zip(old_batch_entries) {
545            let retained_len = entries.iter().filter(|(g, _)| *g >= emit_groups).count();
546
547            if retained_len == 0 {
548                continue;
549            }
550
551            if retained_len == entries.len() {
552                // Nothing was emitted from this batch, so we keep the existing
553                // array and only renumber the remaining group IDs so that they
554                // start from 0.
555                let mut retained_entries = entries;
556                for (g, _) in &mut retained_entries {
557                    *g -= emit_groups;
558                }
559                retained_entries.shrink_to_fit();
560                batches.push(batch);
561                batch_entries.push(retained_entries);
562                continue;
563            }
564
565            let mut retained_entries = Vec::with_capacity(retained_len);
566            let mut retained_rows = Vec::with_capacity(retained_len);
567
568            for (g, r) in entries {
569                if g >= emit_groups {
570                    // Compute the new `(group_idx, row_idx)` pair for a
571                    // retained row. `group_idx` is renumbered to start from
572                    // 0, and `row_idx` points into the new dense batch we are
573                    // building.
574                    retained_entries.push((g - emit_groups, retained_rows.len() as u32));
575                    retained_rows.push(r);
576                }
577            }
578
579            debug_assert_eq!(retained_entries.len(), retained_len);
580            debug_assert_eq!(retained_rows.len(), retained_len);
581
582            let batch = if retained_len == batch.len() {
583                batch
584            } else {
585                // Compact mixed batches so retained rows no longer pin the
586                // original array.
587                let retained_rows = UInt32Array::from(retained_rows);
588                arrow::compute::take(batch.as_ref(), &retained_rows, None)?
589            };
590
591            batches.push(batch);
592            batch_entries.push(retained_entries);
593        }
594
595        self.batches = batches;
596        self.batch_entries = batch_entries;
597        self.num_groups -= emit_groups as usize;
598
599        Ok(())
600    }
601}
602
603impl GroupsAccumulator for ArrayAggGroupsAccumulator {
604    /// Store a reference to the input batch, plus a `(group_idx, row_idx)` pair
605    /// for every row.
606    fn update_batch(
607        &mut self,
608        values: &[ArrayRef],
609        group_indices: &[usize],
610        opt_filter: Option<&BooleanArray>,
611        total_num_groups: usize,
612    ) -> Result<()> {
613        assert_eq!(values.len(), 1, "single argument to update_batch");
614        let input = &values[0];
615
616        self.num_groups = self.num_groups.max(total_num_groups);
617
618        let nulls = if self.ignore_nulls {
619            input.logical_nulls()
620        } else {
621            None
622        };
623
624        let mut entries = Vec::new();
625
626        for (row_idx, &group_idx) in group_indices.iter().enumerate() {
627            // Skip filtered rows
628            if let Some(filter) = opt_filter
629                && (filter.is_null(row_idx) || !filter.value(row_idx))
630            {
631                continue;
632            }
633
634            // Skip null values when ignore_nulls is set
635            if let Some(ref nulls) = nulls
636                && nulls.is_null(row_idx)
637            {
638                continue;
639            }
640
641            entries.push((group_idx as u32, row_idx as u32));
642        }
643
644        // We only need to record the batch if it was non-empty.
645        if !entries.is_empty() {
646            self.batches.push(Arc::clone(input));
647            self.batch_entries.push(entries);
648        }
649
650        Ok(())
651    }
652
653    /// Produce a `ListArray` ordered by group index: the list at
654    /// position N contains the aggregated values for group N.
655    ///
656    /// Uses a counting sort to rearrange the stored `(group, row)`
657    /// entries into group order, then calls `interleave` to gather
658    /// the values into a flat array that backs the output `ListArray`.
659    fn evaluate(&mut self, emit_to: EmitTo) -> Result<ArrayRef> {
660        let emit_groups = match emit_to {
661            EmitTo::All => self.num_groups,
662            EmitTo::First(n) => n,
663        };
664
665        // Step 1: Count entries per group. For EmitTo::First(n), only groups
666        // 0..n are counted; the rest are retained to be emitted in the future.
667        let mut counts = vec![0u32; emit_groups];
668        for entries in &self.batch_entries {
669            for &(g, _) in entries {
670                let g = g as usize;
671                if g < emit_groups {
672                    counts[g] += 1;
673                }
674            }
675        }
676
677        // Step 2: Do a prefix sum over the counts and use it to build ListArray
678        // offsets, null buffer, and write positions for the counting sort.
679        let mut offsets = Vec::<i32>::with_capacity(emit_groups + 1);
680        offsets.push(0);
681        let mut nulls_builder = NullBufferBuilder::new(emit_groups);
682        let mut write_positions = Vec::with_capacity(emit_groups);
683        let mut cur_offset = 0u32;
684        for &count in &counts {
685            if count == 0 {
686                nulls_builder.append_null();
687            } else {
688                nulls_builder.append_non_null();
689            }
690            write_positions.push(cur_offset);
691            cur_offset += count;
692            offsets.push(cur_offset as i32);
693        }
694        let total_rows = cur_offset as usize;
695
696        // Step 3: Scatter entries into group order using the counting sort. The
697        // batch index is implicit from the outer loop position.
698        let flat_values = if total_rows == 0 {
699            new_empty_array(&self.datatype)
700        } else {
701            let mut interleave_indices = vec![(0usize, 0usize); total_rows];
702            for (batch_idx, entries) in self.batch_entries.iter().enumerate() {
703                for &(g, r) in entries {
704                    let g = g as usize;
705                    if g < emit_groups {
706                        let wp = write_positions[g] as usize;
707                        interleave_indices[wp] = (batch_idx, r as usize);
708                        write_positions[g] += 1;
709                    }
710                }
711            }
712
713            let sources: Vec<&dyn Array> =
714                self.batches.iter().map(|b| b.as_ref()).collect();
715            arrow::compute::interleave(&sources, &interleave_indices)?
716        };
717
718        // Step 4: Release state for emitted groups.
719        match emit_to {
720            EmitTo::All => self.clear_state(),
721            EmitTo::First(_) => self.compact_retained_state(emit_groups)?,
722        }
723
724        let offsets = OffsetBuffer::new(ScalarBuffer::from(offsets));
725        let field = Arc::new(Field::new_list_field(self.datatype.clone(), true));
726        let result = ListArray::new(field, offsets, flat_values, nulls_builder.finish());
727
728        Ok(Arc::new(result))
729    }
730
731    fn state(&mut self, emit_to: EmitTo) -> Result<Vec<ArrayRef>> {
732        Ok(vec![self.evaluate(emit_to)?])
733    }
734
735    fn merge_batch(
736        &mut self,
737        values: &[ArrayRef],
738        group_indices: &[usize],
739        total_num_groups: usize,
740    ) -> Result<()> {
741        assert_eq!(values.len(), 1, "one argument to merge_batch");
742        let input_list = values[0].as_list::<i32>();
743
744        self.num_groups = self.num_groups.max(total_num_groups);
745
746        // Push the ListArray's backing values array as a single batch.
747        let list_values = input_list.values();
748        let list_offsets = input_list.offsets();
749
750        let mut entries = Vec::new();
751
752        for (row_idx, &group_idx) in group_indices.iter().enumerate() {
753            if input_list.is_null(row_idx) {
754                continue;
755            }
756            let start = list_offsets[row_idx] as u32;
757            let end = list_offsets[row_idx + 1] as u32;
758            for pos in start..end {
759                entries.push((group_idx as u32, pos));
760            }
761        }
762
763        if !entries.is_empty() {
764            self.batches.push(Arc::clone(list_values));
765            self.batch_entries.push(entries);
766        }
767
768        Ok(())
769    }
770
771    fn convert_to_state(
772        &self,
773        values: &[ArrayRef],
774        opt_filter: Option<&BooleanArray>,
775    ) -> Result<Vec<ArrayRef>> {
776        assert_eq!(values.len(), 1, "one argument to convert_to_state");
777
778        let input = &values[0];
779
780        // Each row becomes a 1-element list: offsets are [0, 1, 2, ..., n].
781        let offsets = OffsetBuffer::from_repeated_length(1, input.len());
782
783        // Filtered rows become null list entries, which merge_batch will skip.
784        let filter_nulls = opt_filter.map(filter_to_nulls);
785
786        // With ignore_nulls, null values also become null list entries. Without
787        // ignore_nulls, null values stay as [NULL] so merge_batch retains them.
788        let nulls = if self.ignore_nulls {
789            let logical = input.logical_nulls();
790            NullBuffer::union(filter_nulls.as_ref(), logical.as_ref())
791        } else {
792            filter_nulls
793        };
794
795        let field = Arc::new(Field::new_list_field(self.datatype.clone(), true));
796        let list_array = ListArray::new(field, offsets, Arc::clone(input), nulls);
797
798        Ok(vec![Arc::new(list_array)])
799    }
800    fn size(&self) -> usize {
801        self.batches
802            .iter()
803            .map(|arr| arr.to_data().get_slice_memory_size().unwrap_or_default())
804            .sum::<usize>()
805            + self.batches.capacity() * size_of::<ArrayRef>()
806            + self
807                .batch_entries
808                .iter()
809                .map(|e| e.capacity() * size_of::<(u32, u32)>())
810                .sum::<usize>()
811            + self.batch_entries.capacity() * size_of::<Vec<(u32, u32)>>()
812    }
813}
814
815/// Resources that are allocated lazily on the first `update_batch` call,
816/// once the concrete runtime Arrow type is known.
817///
818/// Grouping all three fields together makes the "either all present or all
819/// absent" invariant explicit in the type system, replacing the scattered
820/// `.expect()` calls that would otherwise be needed.
821#[derive(Debug)]
822struct DistinctState {
823    /// Converts Arrow arrays to/from the comparable row format.
824    converter: RowConverter,
825    /// One owned encoded row per live distinct value, indexed by group index.
826    /// Compacted via swap-remove on eviction so there are never dead slots.
827    group_rows: Vec<OwnedRow>,
828    /// Live refcount per group index. `counts[i]` is how many times the value
829    /// at `group_rows[i]` is currently present in the window frame.
830    counts: Vec<u64>,
831    /// Hash of the encoded row at group index `i`, kept in sync with
832    /// `group_rows` and `counts`. Needed to patch the map on swap-remove
833    /// eviction without re-encoding the moved row.
834    row_hashes: Vec<u64>,
835    /// Temporary buffer for encoding an incoming batch; reused across calls.
836    rows_buffer: Rows,
837}
838
839#[derive(Debug)]
840pub struct DistinctArrayAggAccumulator {
841    /// Lazily allocated on the first `update_batch`; `None` until then.
842    state: Option<DistinctState>,
843    /// Hash table storing `(hash, group_index)`. Only contains live entries
844    /// (those whose count is > 0). Evicted on `retract_batch` when count
845    /// drops to zero.
846    map: HashTable<(u64, usize)>,
847    /// Heap size of `map` in bytes, tracked for `size()` reporting.
848    map_size: usize,
849    /// Reused buffer for batch hashes.
850    hashes_buffer: Vec<u64>,
851    /// Random state used by `create_hashes`.
852    random_state: RandomState,
853    datatype: DataType,
854    sort_options: Option<SortOptions>,
855    ignore_nulls: bool,
856}
857
858/// Returns `true` if `dt` is, or recursively contains, a `Dictionary` type.
859///
860/// `RowConverter` always decodes to the physical (non-dictionary) type, so a
861/// cast back to the declared logical type is required when this is true.
862fn datatype_contains_dictionary(dt: &DataType) -> bool {
863    match dt {
864        DataType::Dictionary(_, _) => true,
865        DataType::List(f)
866        | DataType::LargeList(f)
867        | DataType::FixedSizeList(f, _)
868        | DataType::Map(f, _) => datatype_contains_dictionary(f.data_type()),
869        DataType::Struct(fields) => fields
870            .iter()
871            .any(|f| datatype_contains_dictionary(f.data_type())),
872        _ => false,
873    }
874}
875
876impl DistinctArrayAggAccumulator {
877    pub fn try_new(
878        datatype: &DataType,
879        sort_options: Option<SortOptions>,
880        ignore_nulls: bool,
881    ) -> Result<Self> {
882        Ok(Self {
883            state: None,
884            map: HashTable::new(),
885            map_size: 0,
886            hashes_buffer: Vec::new(),
887            random_state: RandomState::default(),
888            datatype: datatype.clone(),
889            sort_options,
890            ignore_nulls,
891        })
892    }
893
894    /// Lazily initialises the `DistinctState` on the first call, using the
895    /// actual runtime column type.
896    fn ensure_state(&mut self, data_type: &DataType) -> Result<()> {
897        if self.state.is_none() {
898            let sort_field = match self.sort_options {
899                Some(opts) => SortField::new_with_options(data_type.clone(), opts),
900                None => SortField::new(data_type.clone()),
901            };
902            let converter = RowConverter::new(vec![sort_field])?;
903            let rows_buffer = converter.empty_rows(0, 0);
904            self.state = Some(DistinctState {
905                converter,
906                group_rows: Vec::new(),
907                counts: Vec::new(),
908                row_hashes: Vec::new(),
909                rows_buffer,
910            });
911        }
912        Ok(())
913    }
914}
915
916impl Accumulator for DistinctArrayAggAccumulator {
917    fn state(&mut self) -> Result<Vec<ScalarValue>> {
918        Ok(vec![self.evaluate()?])
919    }
920
921    fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
922        if values.is_empty() {
923            return Ok(());
924        }
925
926        let val = &values[0];
927
928        // Filter nulls out upfront when ignore_nulls is set so they are
929        // never inserted into the dedup state.
930        let filtered;
931        let col: &ArrayRef = if self.ignore_nulls {
932            if let Some(nulls) = val.logical_nulls() {
933                if nulls.null_count() > 0 {
934                    let mask: BooleanArray = nulls.iter().map(Some).collect();
935                    filtered = filter(val.as_ref(), &mask)?;
936                    &filtered
937                } else {
938                    val
939                }
940            } else {
941                val
942            }
943        } else {
944            val
945        };
946
947        if col.is_empty() {
948            return Ok(());
949        }
950
951        self.ensure_state(col.data_type())?;
952
953        // Encode the entire incoming batch into rows_buffer in one pass.
954        let DistinctState {
955            converter,
956            group_rows,
957            counts,
958            row_hashes,
959            rows_buffer,
960        } = self.state.as_mut().unwrap();
961        rows_buffer.clear();
962        converter.append(rows_buffer, std::slice::from_ref(col))?;
963
964        // Pre-compute all hashes for the batch in one SIMD-friendly pass.
965        self.hashes_buffer.clear();
966        self.hashes_buffer.resize(col.len(), 0);
967        create_hashes(
968            std::slice::from_ref(col),
969            &self.random_state,
970            &mut self.hashes_buffer,
971        )?;
972
973        for (row_idx, &hash) in self.hashes_buffer.iter().enumerate() {
974            let row = rows_buffer.row(row_idx);
975            let entry = self.map.find_mut(hash, |&(h, group_idx)| {
976                h == hash && group_rows[group_idx].row() == row
977            });
978            match entry {
979                Some((_, group_idx)) => {
980                    // Already known: just increment the live refcount.
981                    counts[*group_idx] += 1;
982                }
983                None => {
984                    // New distinct value: own the encoded row, record it.
985                    let new_group_idx = group_rows.len();
986                    group_rows.push(row.owned());
987                    counts.push(1);
988                    row_hashes.push(hash);
989                    self.map.insert_accounted(
990                        (hash, new_group_idx),
991                        |&(h, _)| h,
992                        &mut self.map_size,
993                    );
994                }
995            }
996        }
997        Ok(())
998    }
999
1000    fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
1001        if states.is_empty() {
1002            return Ok(());
1003        }
1004
1005        assert_eq_or_internal_err!(states.len(), 1, "expects single state");
1006
1007        // The DISTINCT state is `List<value>`.
1008        states[0]
1009            .as_list::<i32>()
1010            .iter()
1011            .flatten()
1012            .try_for_each(|val| self.update_batch(&[val]))
1013    }
1014
1015    fn evaluate(&mut self) -> Result<ScalarValue> {
1016        if self.map.is_empty() {
1017            return Ok(ScalarValue::new_null_list(self.datatype.clone(), true, 1));
1018        }
1019
1020        let DistinctState {
1021            converter,
1022            group_rows,
1023            ..
1024        } = self
1025            .state
1026            .as_ref()
1027            .expect("state must be set when map is non-empty");
1028
1029        // Collect the group indices of all live entries.
1030        let mut live_indices: Vec<usize> =
1031            self.map.iter().map(|&(_, group_idx)| group_idx).collect();
1032
1033        // If ORDER BY was specified, the RowConverter bakes the sort direction
1034        // into the row bytes, so lexicographic sort gives the correct order.
1035        if self.sort_options.is_some() {
1036            live_indices
1037                .sort_unstable_by(|&a, &b| group_rows[a].row().cmp(&group_rows[b].row()));
1038        }
1039
1040        // Decode the selected rows back into an Arrow array.
1041        let rows: Vec<Row<'_>> =
1042            live_indices.iter().map(|&i| group_rows[i].row()).collect();
1043        let arrays = converter.convert_rows(rows)?;
1044
1045        // `convert_rows` always returns the physical (non-dictionary) type.
1046        // Cast back to the declared logical type when they differ AND the
1047        // declared type contains a Dictionary somewhere (directly or nested
1048        // inside a Struct, List, etc.) — that is the only case where
1049        // RowConverter strips the logical type.
1050        let decoded = if arrays[0].data_type() != &self.datatype
1051            && datatype_contains_dictionary(&self.datatype)
1052        {
1053            cast(arrays[0].as_ref(), &self.datatype)?
1054        } else {
1055            Arc::clone(&arrays[0])
1056        };
1057
1058        let values: Vec<ScalarValue> = (0..decoded.len())
1059            .map(|i| ScalarValue::try_from_array(decoded.as_ref(), i))
1060            .collect::<Result<_>>()?;
1061
1062        let arr = ScalarValue::new_list(&values, &self.datatype, true);
1063        Ok(ScalarValue::List(arr))
1064    }
1065
1066    fn retract_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
1067        if values.is_empty() {
1068            return Ok(());
1069        }
1070
1071        assert_eq_or_internal_err!(values.len(), 1, "expects single batch");
1072
1073        let val = &values[0];
1074
1075        // Mirror the null-filtering logic from update_batch so we only
1076        // retract values that were actually inserted.
1077        let filtered;
1078        let col: &ArrayRef = if self.ignore_nulls {
1079            if let Some(nulls) = val.logical_nulls() {
1080                if nulls.null_count() > 0 {
1081                    let mask: BooleanArray = nulls.iter().map(Some).collect();
1082                    filtered = filter(val.as_ref(), &mask)?;
1083                    &filtered
1084                } else {
1085                    val
1086                }
1087            } else {
1088                val
1089            }
1090        } else {
1091            val
1092        };
1093
1094        if col.is_empty() {
1095            return Ok(());
1096        }
1097
1098        let DistinctState {
1099            converter,
1100            group_rows,
1101            counts,
1102            row_hashes,
1103            rows_buffer,
1104        } = self
1105            .state
1106            .as_mut()
1107            .expect("retract_batch called before update_batch");
1108
1109        rows_buffer.clear();
1110        converter.append(rows_buffer, std::slice::from_ref(col))?;
1111
1112        self.hashes_buffer.clear();
1113        self.hashes_buffer.resize(col.len(), 0);
1114        create_hashes(
1115            std::slice::from_ref(col),
1116            &self.random_state,
1117            &mut self.hashes_buffer,
1118        )?;
1119
1120        for (row_idx, &hash) in self.hashes_buffer.iter().enumerate() {
1121            let row = rows_buffer.row(row_idx);
1122            match self.map.find_entry(hash, |&(h, group_idx)| {
1123                h == hash && group_rows[group_idx].row() == row
1124            }) {
1125                Err(_) => {
1126                    return internal_err!(
1127                        "DistinctArrayAggAccumulator::retract_batch: \
1128                         value not present in state"
1129                    );
1130                }
1131                Ok(occupied) => {
1132                    let (_, dead_idx) = *occupied.get();
1133                    counts[dead_idx] -= 1;
1134                    if counts[dead_idx] == 0 {
1135                        occupied.remove();
1136                        // Compact via swap-remove: move the last slot into the
1137                        // dead slot so group_rows / counts / row_hashes stay
1138                        // dense with no dead entries.
1139                        let last_idx = group_rows.len() - 1;
1140                        if dead_idx != last_idx {
1141                            // Patch the map entry that points to last_idx so
1142                            // it points to dead_idx instead.
1143                            let last_hash = row_hashes[last_idx];
1144                            self.map
1145                                .find_mut(last_hash, |&(_, idx)| idx == last_idx)
1146                                .ok_or_else(|| {
1147                                    datafusion_common::internal_datafusion_err!(
1148                                        "DistinctArrayAggAccumulator: map is missing \
1149                                         group index {last_idx} during swap-remove \
1150                                         compaction"
1151                                    )
1152                                })?
1153                                .1 = dead_idx;
1154                        }
1155                        group_rows.swap_remove(dead_idx);
1156                        counts.swap_remove(dead_idx);
1157                        row_hashes.swap_remove(dead_idx);
1158                    }
1159                }
1160            }
1161        }
1162        Ok(())
1163    }
1164
1165    fn supports_retract_batch(&self) -> bool {
1166        true
1167    }
1168
1169    fn size(&self) -> usize {
1170        size_of_val(self)
1171            + self
1172                .state
1173                .as_ref()
1174                .map(|s| {
1175                    s.group_rows
1176                        .iter()
1177                        .map(|r| r.row().data().len())
1178                        .sum::<usize>()
1179                        + s.group_rows.capacity() * size_of::<OwnedRow>()
1180                        + s.counts.capacity() * size_of::<u64>()
1181                        + s.row_hashes.capacity() * size_of::<u64>()
1182                        + s.rows_buffer.size()
1183                        + s.converter.size()
1184                })
1185                .unwrap_or(0)
1186            + self.map_size
1187            + self.hashes_buffer.capacity() * size_of::<u64>()
1188            + self.datatype.size()
1189            - size_of_val(&self.datatype)
1190    }
1191}
1192
1193/// Accumulator for a `ARRAY_AGG(... ORDER BY ..., ...)` aggregation. In a multi
1194/// partition setting, partial aggregations are computed for every partition,
1195/// and then their results are merged.
1196#[derive(Debug)]
1197pub(crate) struct OrderSensitiveArrayAggAccumulator {
1198    /// Stores entries in the `ARRAY_AGG` result.
1199    values: Vec<ScalarValue>,
1200    /// Stores values of ordering requirement expressions corresponding to each
1201    /// entry in `values`. This information is used when merging results from
1202    /// different partitions. For detailed information how merging is done, see
1203    /// [`merge_ordered_arrays`].
1204    ordering_values: Vec<Vec<ScalarValue>>,
1205    /// Stores datatypes of expressions inside values and ordering requirement
1206    /// expressions.
1207    datatypes: Vec<DataType>,
1208    /// Stores the ordering requirement of the `Accumulator`.
1209    ordering_req: LexOrdering,
1210    /// Whether the input is known to be pre-ordered
1211    is_input_pre_ordered: bool,
1212    /// Whether the aggregation is running in reverse.
1213    reverse: bool,
1214    /// Whether the aggregation should ignore null values.
1215    ignore_nulls: bool,
1216}
1217
1218impl OrderSensitiveArrayAggAccumulator {
1219    /// Create a new order-sensitive ARRAY_AGG accumulator based on the given
1220    /// item data type.
1221    pub fn try_new(
1222        datatype: &DataType,
1223        ordering_dtypes: &[DataType],
1224        ordering_req: LexOrdering,
1225        is_input_pre_ordered: bool,
1226        reverse: bool,
1227        ignore_nulls: bool,
1228    ) -> Result<Self> {
1229        let mut datatypes = vec![datatype.clone()];
1230        datatypes.extend(ordering_dtypes.iter().cloned());
1231        Ok(Self {
1232            values: vec![],
1233            ordering_values: vec![],
1234            datatypes,
1235            ordering_req,
1236            is_input_pre_ordered,
1237            reverse,
1238            ignore_nulls,
1239        })
1240    }
1241
1242    fn sort(&mut self) {
1243        let sort_options = self
1244            .ordering_req
1245            .iter()
1246            .map(|sort_expr| sort_expr.options)
1247            .collect::<Vec<_>>();
1248        let mut values = take(&mut self.values)
1249            .into_iter()
1250            .zip(take(&mut self.ordering_values))
1251            .collect::<Vec<_>>();
1252        let mut delayed_cmp_err = Ok(());
1253        values.sort_by(|(_, left_ordering), (_, right_ordering)| {
1254            compare_rows(left_ordering, right_ordering, &sort_options).unwrap_or_else(
1255                |err| {
1256                    delayed_cmp_err = Err(err);
1257                    Ordering::Equal
1258                },
1259            )
1260        });
1261        (self.values, self.ordering_values) = values.into_iter().unzip();
1262    }
1263
1264    fn evaluate_orderings(&self) -> Result<ScalarValue> {
1265        let fields = ordering_fields(&self.ordering_req, &self.datatypes[1..]);
1266
1267        let column_wise_ordering_values = if self.ordering_values.is_empty() {
1268            fields
1269                .iter()
1270                .map(|f| new_empty_array(f.data_type()))
1271                .collect::<Vec<_>>()
1272        } else {
1273            (0..fields.len())
1274                .map(|i| {
1275                    let column_values: Box<dyn Iterator<Item = ScalarValue>> = if self
1276                        .reverse
1277                    {
1278                        Box::new(self.ordering_values.iter().rev().map(|x| x[i].clone()))
1279                    } else {
1280                        Box::new(self.ordering_values.iter().map(|x| x[i].clone()))
1281                    };
1282                    ScalarValue::iter_to_array(column_values)
1283                })
1284                .collect::<Result<_>>()?
1285        };
1286
1287        let ordering_array = StructArray::try_new(
1288            Fields::from(fields),
1289            column_wise_ordering_values,
1290            None,
1291        )?;
1292        Ok(SingleRowListArrayBuilder::new(Arc::new(ordering_array)).build_list_scalar())
1293    }
1294}
1295
1296impl Accumulator for OrderSensitiveArrayAggAccumulator {
1297    fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
1298        if values.is_empty() {
1299            return Ok(());
1300        }
1301
1302        let val = &values[0];
1303        let ord = &values[1..];
1304        let nulls = if self.ignore_nulls {
1305            val.logical_nulls()
1306        } else {
1307            None
1308        };
1309
1310        let nulls = nulls.as_ref();
1311        if nulls.is_none_or(|nulls| nulls.null_count() < val.len()) {
1312            for i in 0..val.len() {
1313                if nulls.is_none_or(|nulls| nulls.is_valid(i)) {
1314                    self.values
1315                        .push(ScalarValue::try_from_array(val, i)?.compacted());
1316                    self.ordering_values.push(
1317                        get_row_at_idx(ord, i)?
1318                            .into_iter()
1319                            .map(|v| v.compacted())
1320                            .collect(),
1321                    )
1322                }
1323            }
1324        }
1325
1326        Ok(())
1327    }
1328
1329    fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
1330        if states.is_empty() {
1331            return Ok(());
1332        }
1333
1334        // First entry in the state is the aggregation result. Second entry
1335        // stores values received for ordering requirement columns for each
1336        // aggregation value inside `ARRAY_AGG` list. For each `StructArray`
1337        // inside `ARRAY_AGG` list, we will receive an `Array` that stores values
1338        // received from its ordering requirement expression. (This information
1339        // is necessary for during merging).
1340        let [array_agg_values, agg_orderings] =
1341            take_function_args("OrderSensitiveArrayAggAccumulator::merge_batch", states)?;
1342        let Some(agg_orderings) = agg_orderings.as_list_opt::<i32>() else {
1343            return exec_err!("Expects to receive a list array");
1344        };
1345
1346        // Stores ARRAY_AGG results coming from each partition
1347        let mut partition_values = vec![];
1348        // Stores ordering requirement expression results coming from each partition
1349        let mut partition_ordering_values = vec![];
1350
1351        // Existing values should be merged also.
1352        if !self.is_input_pre_ordered {
1353            self.sort();
1354        }
1355        partition_values.push(take(&mut self.values).into());
1356        partition_ordering_values.push(take(&mut self.ordering_values).into());
1357
1358        // Convert array to Scalars to sort them easily. Convert back to array at evaluation.
1359        let array_agg_res = ScalarValue::convert_array_to_scalar_vec(array_agg_values)?;
1360        for maybe_v in array_agg_res.into_iter() {
1361            if let Some(v) = maybe_v {
1362                partition_values.push(v.into());
1363            } else {
1364                partition_values.push(vec![].into());
1365            }
1366        }
1367
1368        let orderings = ScalarValue::convert_array_to_scalar_vec(agg_orderings)?;
1369        for partition_ordering_rows in orderings.into_iter().flatten() {
1370            // Extract value from struct to ordering_rows for each group/partition
1371            let ordering_value = partition_ordering_rows.into_iter().map(|ordering_row| {
1372                    if let ScalarValue::Struct(s) = ordering_row {
1373                        let mut ordering_columns_per_row = vec![];
1374
1375                        for column in s.columns() {
1376                            let sv = ScalarValue::try_from_array(column, 0)?;
1377                            ordering_columns_per_row.push(sv);
1378                        }
1379
1380                        Ok(ordering_columns_per_row)
1381                    } else {
1382                        exec_err!(
1383                            "Expects to receive ScalarValue::Struct(Arc<StructArray>) but got:{:?}",
1384                            ordering_row.data_type()
1385                        )
1386                    }
1387                }).collect::<Result<VecDeque<_>>>()?;
1388
1389            partition_ordering_values.push(ordering_value);
1390        }
1391
1392        let sort_options = self
1393            .ordering_req
1394            .iter()
1395            .map(|sort_expr| sort_expr.options)
1396            .collect::<Vec<_>>();
1397
1398        (self.values, self.ordering_values) = merge_ordered_arrays(
1399            &mut partition_values,
1400            &mut partition_ordering_values,
1401            &sort_options,
1402        )?;
1403
1404        Ok(())
1405    }
1406
1407    fn state(&mut self) -> Result<Vec<ScalarValue>> {
1408        if !self.is_input_pre_ordered {
1409            self.sort();
1410        }
1411
1412        let mut result = vec![self.evaluate()?];
1413        result.push(self.evaluate_orderings()?);
1414
1415        Ok(result)
1416    }
1417
1418    fn evaluate(&mut self) -> Result<ScalarValue> {
1419        if !self.is_input_pre_ordered {
1420            self.sort();
1421        }
1422
1423        if self.values.is_empty() {
1424            return Ok(ScalarValue::new_null_list(
1425                self.datatypes[0].clone(),
1426                true,
1427                1,
1428            ));
1429        }
1430
1431        let values = self.values.clone();
1432        let array = if self.reverse {
1433            ScalarValue::new_list_from_iter(
1434                values.into_iter().rev(),
1435                &self.datatypes[0],
1436                true,
1437            )
1438        } else {
1439            ScalarValue::new_list_from_iter(values.into_iter(), &self.datatypes[0], true)
1440        };
1441        Ok(ScalarValue::List(array))
1442    }
1443
1444    fn size(&self) -> usize {
1445        let mut total = size_of_val(self) + ScalarValue::size_of_vec(&self.values)
1446            - size_of_val(&self.values);
1447
1448        // Add size of the `self.ordering_values`
1449        total += size_of::<Vec<ScalarValue>>() * self.ordering_values.capacity();
1450        for row in &self.ordering_values {
1451            total += ScalarValue::size_of_vec(row) - size_of_val(row);
1452        }
1453
1454        // Add size of the `self.datatypes`
1455        total += size_of::<DataType>() * self.datatypes.capacity();
1456        for dtype in &self.datatypes {
1457            total += dtype.size() - size_of_val(dtype);
1458        }
1459
1460        // Add size of the `self.ordering_req`
1461        total += size_of::<PhysicalSortExpr>() * self.ordering_req.capacity();
1462        // TODO: Calculate size of each `PhysicalSortExpr` more accurately.
1463        total
1464    }
1465}
1466
1467#[cfg(test)]
1468mod tests {
1469    use super::*;
1470    use arrow::array::{ListBuilder, StringBuilder};
1471    use arrow::datatypes::Schema;
1472    use datafusion_common::cast::as_generic_string_array;
1473    use datafusion_common::internal_err;
1474    use datafusion_physical_expr::PhysicalExpr;
1475    use datafusion_physical_expr::expressions::Column;
1476
1477    #[test]
1478    fn no_duplicates_no_distinct() -> Result<()> {
1479        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string().build_two()?;
1480
1481        acc1.update_batch(&[data(["a", "b", "c"])])?;
1482        acc2.update_batch(&[data(["d", "e", "f"])])?;
1483        acc1 = merge(acc1, acc2)?;
1484
1485        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1486
1487        assert_eq!(result, vec!["a", "b", "c", "d", "e", "f"]);
1488
1489        Ok(())
1490    }
1491
1492    #[test]
1493    fn no_duplicates_distinct() -> Result<()> {
1494        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1495            .distinct()
1496            .build_two()?;
1497
1498        acc1.update_batch(&[data(["a", "b", "c"])])?;
1499        acc2.update_batch(&[data(["d", "e", "f"])])?;
1500        acc1 = merge(acc1, acc2)?;
1501
1502        let mut result = print_nulls(str_arr(acc1.evaluate()?)?);
1503        result.sort();
1504
1505        assert_eq!(result, vec!["a", "b", "c", "d", "e", "f"]);
1506
1507        Ok(())
1508    }
1509
1510    #[test]
1511    fn duplicates_no_distinct() -> Result<()> {
1512        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string().build_two()?;
1513
1514        acc1.update_batch(&[data(["a", "b", "c"])])?;
1515        acc2.update_batch(&[data(["a", "b", "c"])])?;
1516        acc1 = merge(acc1, acc2)?;
1517
1518        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1519
1520        assert_eq!(result, vec!["a", "b", "c", "a", "b", "c"]);
1521
1522        Ok(())
1523    }
1524
1525    #[test]
1526    fn duplicates_distinct() -> Result<()> {
1527        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1528            .distinct()
1529            .build_two()?;
1530
1531        acc1.update_batch(&[data(["a", "b", "c"])])?;
1532        acc2.update_batch(&[data(["a", "b", "c"])])?;
1533        acc1 = merge(acc1, acc2)?;
1534
1535        let mut result = print_nulls(str_arr(acc1.evaluate()?)?);
1536        result.sort();
1537
1538        assert_eq!(result, vec!["a", "b", "c"]);
1539
1540        Ok(())
1541    }
1542
1543    #[test]
1544    fn duplicates_on_second_batch_distinct() -> Result<()> {
1545        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1546            .distinct()
1547            .build_two()?;
1548
1549        acc1.update_batch(&[data(["a", "c"])])?;
1550        acc2.update_batch(&[data(["d", "a", "b", "c"])])?;
1551        acc1 = merge(acc1, acc2)?;
1552
1553        let mut result = print_nulls(str_arr(acc1.evaluate()?)?);
1554        result.sort();
1555
1556        assert_eq!(result, vec!["a", "b", "c", "d"]);
1557
1558        Ok(())
1559    }
1560
1561    #[test]
1562    fn no_duplicates_distinct_sort_asc() -> Result<()> {
1563        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1564            .distinct()
1565            .order_by_col("col", SortOptions::new(false, false))
1566            .build_two()?;
1567
1568        acc1.update_batch(&[data(["e", "b", "d"])])?;
1569        acc2.update_batch(&[data(["f", "a", "c"])])?;
1570        acc1 = merge(acc1, acc2)?;
1571
1572        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1573
1574        assert_eq!(result, vec!["a", "b", "c", "d", "e", "f"]);
1575
1576        Ok(())
1577    }
1578
1579    #[test]
1580    fn no_duplicates_distinct_sort_desc() -> Result<()> {
1581        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1582            .distinct()
1583            .order_by_col("col", SortOptions::new(true, false))
1584            .build_two()?;
1585
1586        acc1.update_batch(&[data(["e", "b", "d"])])?;
1587        acc2.update_batch(&[data(["f", "a", "c"])])?;
1588        acc1 = merge(acc1, acc2)?;
1589
1590        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1591
1592        assert_eq!(result, vec!["f", "e", "d", "c", "b", "a"]);
1593
1594        Ok(())
1595    }
1596
1597    #[test]
1598    fn duplicates_distinct_sort_asc() -> Result<()> {
1599        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1600            .distinct()
1601            .order_by_col("col", SortOptions::new(false, false))
1602            .build_two()?;
1603
1604        acc1.update_batch(&[data(["a", "c", "b"])])?;
1605        acc2.update_batch(&[data(["b", "c", "a"])])?;
1606        acc1 = merge(acc1, acc2)?;
1607
1608        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1609
1610        assert_eq!(result, vec!["a", "b", "c"]);
1611
1612        Ok(())
1613    }
1614
1615    #[test]
1616    fn duplicates_distinct_sort_desc() -> Result<()> {
1617        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1618            .distinct()
1619            .order_by_col("col", SortOptions::new(true, false))
1620            .build_two()?;
1621
1622        acc1.update_batch(&[data(["a", "c", "b"])])?;
1623        acc2.update_batch(&[data(["b", "c", "a"])])?;
1624        acc1 = merge(acc1, acc2)?;
1625
1626        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1627
1628        assert_eq!(result, vec!["c", "b", "a"]);
1629
1630        Ok(())
1631    }
1632
1633    #[test]
1634    fn no_duplicates_distinct_sort_asc_nulls_first() -> Result<()> {
1635        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1636            .distinct()
1637            .order_by_col("col", SortOptions::new(false, true))
1638            .build_two()?;
1639
1640        acc1.update_batch(&[data([Some("e"), Some("b"), None])])?;
1641        acc2.update_batch(&[data([Some("f"), Some("a"), None])])?;
1642        acc1 = merge(acc1, acc2)?;
1643
1644        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1645
1646        assert_eq!(result, vec!["NULL", "a", "b", "e", "f"]);
1647
1648        Ok(())
1649    }
1650
1651    #[test]
1652    fn no_duplicates_distinct_sort_asc_nulls_last() -> Result<()> {
1653        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1654            .distinct()
1655            .order_by_col("col", SortOptions::new(false, false))
1656            .build_two()?;
1657
1658        acc1.update_batch(&[data([Some("e"), Some("b"), None])])?;
1659        acc2.update_batch(&[data([Some("f"), Some("a"), None])])?;
1660        acc1 = merge(acc1, acc2)?;
1661
1662        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1663
1664        assert_eq!(result, vec!["a", "b", "e", "f", "NULL"]);
1665
1666        Ok(())
1667    }
1668
1669    #[test]
1670    fn no_duplicates_distinct_sort_desc_nulls_first() -> Result<()> {
1671        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1672            .distinct()
1673            .order_by_col("col", SortOptions::new(true, true))
1674            .build_two()?;
1675
1676        acc1.update_batch(&[data([Some("e"), Some("b"), None])])?;
1677        acc2.update_batch(&[data([Some("f"), Some("a"), None])])?;
1678        acc1 = merge(acc1, acc2)?;
1679
1680        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1681
1682        assert_eq!(result, vec!["NULL", "f", "e", "b", "a"]);
1683
1684        Ok(())
1685    }
1686
1687    #[test]
1688    fn no_duplicates_distinct_sort_desc_nulls_last() -> Result<()> {
1689        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1690            .distinct()
1691            .order_by_col("col", SortOptions::new(true, false))
1692            .build_two()?;
1693
1694        acc1.update_batch(&[data([Some("e"), Some("b"), None])])?;
1695        acc2.update_batch(&[data([Some("f"), Some("a"), None])])?;
1696        acc1 = merge(acc1, acc2)?;
1697
1698        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1699
1700        assert_eq!(result, vec!["f", "e", "b", "a", "NULL"]);
1701
1702        Ok(())
1703    }
1704
1705    #[test]
1706    fn all_nulls_on_first_batch_with_distinct() -> Result<()> {
1707        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1708            .distinct()
1709            .build_two()?;
1710
1711        acc1.update_batch(&[data::<Option<&str>, 3>([None, None, None])])?;
1712        acc2.update_batch(&[data([Some("a"), None, None, None])])?;
1713        acc1 = merge(acc1, acc2)?;
1714
1715        let mut result = print_nulls(str_arr(acc1.evaluate()?)?);
1716        result.sort();
1717        assert_eq!(result, vec!["NULL", "a"]);
1718        Ok(())
1719    }
1720
1721    #[test]
1722    fn all_nulls_on_both_batches_with_distinct() -> Result<()> {
1723        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string()
1724            .distinct()
1725            .build_two()?;
1726
1727        acc1.update_batch(&[data::<Option<&str>, 3>([None, None, None])])?;
1728        acc2.update_batch(&[data::<Option<&str>, 4>([None, None, None, None])])?;
1729        acc1 = merge(acc1, acc2)?;
1730
1731        let result = print_nulls(str_arr(acc1.evaluate()?)?);
1732        assert_eq!(result, vec!["NULL"]);
1733        Ok(())
1734    }
1735
1736    #[test]
1737    fn does_not_over_account_memory() -> Result<()> {
1738        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::string().build_two()?;
1739
1740        acc1.update_batch(&[data(["a", "c", "b"])])?;
1741        acc2.update_batch(&[data(["b", "c", "a"])])?;
1742        acc1 = merge(acc1, acc2)?;
1743
1744        assert_eq!(acc1.size(), 174);
1745
1746        Ok(())
1747    }
1748    #[test]
1749    fn does_not_over_account_memory_distinct() -> Result<()> {
1750        let (mut acc1, mut acc2) = ArrayAggAccumulatorBuilder::new(DataType::List(
1751            Arc::new(Field::new_list_field(DataType::Utf8, true)),
1752        ))
1753        .distinct()
1754        .build_two()?;
1755
1756        acc1.update_batch(&[string_list_data([
1757            vec!["a", "b", "c"],
1758            vec!["d", "e", "f"],
1759        ])])?;
1760        acc2.update_batch(&[string_list_data([vec!["e", "f", "g"]])])?;
1761        acc1 = merge(acc1, acc2)?;
1762
1763        assert_eq!(acc1.size(), 2274);
1764
1765        Ok(())
1766    }
1767
1768    #[test]
1769    fn does_not_over_account_memory_ordered() -> Result<()> {
1770        let mut acc = ArrayAggAccumulatorBuilder::new(DataType::List(Arc::new(
1771            Field::new_list_field(DataType::Utf8, true),
1772        )))
1773        .order_by_col("col", SortOptions::new(false, false))
1774        .build()?;
1775
1776        acc.update_batch(&[string_list_data([
1777            vec!["a", "b", "c"],
1778            vec!["c", "d", "e"],
1779            vec!["b", "c", "d"],
1780        ])])?;
1781
1782        // without compaction, the size is 17112
1783        assert_eq!(acc.size(), 2224);
1784
1785        Ok(())
1786    }
1787
1788    #[test]
1789    fn ordered_aggregate_nested_nullability_mismatch_issue_24022() -> Result<()> {
1790        use arrow::array::{Int32Array, Int64Array, StructArray};
1791        use datafusion_physical_expr::expressions::Column;
1792
1793        let requested_element_type =
1794            DataType::Struct(Fields::from(vec![Field::new("n", DataType::Int32, true)]));
1795        let inferred_field = Field::new("n", DataType::Int32, false);
1796
1797        let ordering_dtype = DataType::Int64;
1798        let schema = Schema::new(vec![
1799            Field::new("val", requested_element_type.clone(), true),
1800            Field::new("ord", DataType::Int64, true),
1801        ]);
1802        let ord_expr = Arc::new(
1803            Column::new_with_schema("ord", &schema).expect("column not in schema"),
1804        ) as Arc<dyn PhysicalExpr>;
1805
1806        let asc_opts = SortOptions {
1807            descending: false,
1808            nulls_first: false,
1809        };
1810        let asc_ordering = LexOrdering::new(vec![PhysicalSortExpr::new(
1811            Arc::clone(&ord_expr),
1812            asc_opts,
1813        )])
1814        .unwrap();
1815
1816        let mut acc = OrderSensitiveArrayAggAccumulator::try_new(
1817            &requested_element_type,
1818            std::slice::from_ref(&ordering_dtype),
1819            asc_ordering,
1820            /*is_input_pre_ordered=*/ true,
1821            /*reverse=*/ false,
1822            /*ignore_nulls=*/ false,
1823        )?;
1824
1825        let value_arr = Arc::new(StructArray::from(vec![(
1826            Arc::new(inferred_field),
1827            Arc::new(Int32Array::from(vec![1])) as ArrayRef,
1828        )])) as ArrayRef;
1829
1830        let ord_arr = Arc::new(Int64Array::from(vec![0i64])) as ArrayRef;
1831
1832        acc.update_batch(&[value_arr, ord_arr])?;
1833
1834        let evaluated = acc.evaluate()?;
1835
1836        if let ScalarValue::List(arr) = evaluated {
1837            assert_eq!(
1838                arr.data_type(),
1839                &DataType::List(Arc::new(Field::new_list_field(
1840                    requested_element_type.clone(),
1841                    true
1842                )))
1843            );
1844
1845            let expected_struct_array = StructArray::from(vec![(
1846                Arc::new(Field::new("n", DataType::Int32, true)),
1847                Arc::new(Int32Array::from(vec![1])) as ArrayRef,
1848            )]);
1849            let expected_array = Arc::new(expected_struct_array) as ArrayRef;
1850            assert_eq!(&arr.value(0), &expected_array);
1851        } else {
1852            panic!("Expected ScalarValue::List");
1853        }
1854
1855        Ok(())
1856    }
1857
1858    #[test]
1859    fn distinct_aggregate_nested_nullability_mismatch_issue_24022() -> Result<()> {
1860        use arrow::array::{Int32Array, StructArray};
1861        use datafusion_common::ScalarValue;
1862
1863        let requested_element_type =
1864            DataType::Struct(Fields::from(vec![Field::new("n", DataType::Int32, true)]));
1865        let inferred_field = Field::new("n", DataType::Int32, false);
1866
1867        let mut acc = DistinctArrayAggAccumulator::try_new(
1868            &requested_element_type,
1869            None,
1870            /*ignore_nulls=*/ false,
1871        )?;
1872
1873        let value_arr = Arc::new(StructArray::from(vec![(
1874            Arc::new(inferred_field),
1875            Arc::new(Int32Array::from(vec![1])) as ArrayRef,
1876        )])) as ArrayRef;
1877
1878        acc.update_batch(&[value_arr])?;
1879
1880        let evaluated = acc.evaluate()?;
1881
1882        if let ScalarValue::List(arr) = evaluated {
1883            assert_eq!(
1884                arr.data_type(),
1885                &DataType::List(Arc::new(Field::new_list_field(
1886                    requested_element_type.clone(),
1887                    true
1888                )))
1889            );
1890
1891            let expected_struct_array = StructArray::from(vec![(
1892                Arc::new(Field::new("n", DataType::Int32, true)),
1893                Arc::new(Int32Array::from(vec![1])) as ArrayRef,
1894            )]);
1895            let expected_array = Arc::new(expected_struct_array) as ArrayRef;
1896            assert_eq!(&arr.value(0), &expected_array);
1897        } else {
1898            panic!("Expected ScalarValue::List");
1899        }
1900
1901        Ok(())
1902    }
1903
1904    // Reproduces the bug where `state()` emits reversed values but non-reversed
1905    // orderings when the optimizer sets is_input_pre_ordered=true + reverse=true
1906    // (DESC aggregate with ASC pre-sorted input). The partial states are fed into
1907    // a final accumulator via merge_batch; without the fix the ordering keys and
1908    // values are mismatched so the final sort produces wrong order.
1909    #[test]
1910    fn desc_order_partial_final_merge_correct() -> Result<()> {
1911        use arrow::array::Int64Array;
1912        use datafusion_physical_expr::expressions::Column;
1913
1914        let schema = Schema::new(vec![
1915            Field::new("val", DataType::Int64, true),
1916            Field::new("ord", DataType::Int64, true),
1917        ]);
1918        let ord_expr = Arc::new(
1919            Column::new_with_schema("ord", &schema).expect("column not in schema"),
1920        ) as Arc<dyn PhysicalExpr>;
1921
1922        // ordering_req for partial = [ord ASC] (reversed, because input is pre-sorted ASC
1923        // and the user wants DESC — the optimizer reverses the requirement)
1924        let asc_opts = SortOptions {
1925            descending: false,
1926            nulls_first: false,
1927        };
1928        let desc_opts = SortOptions {
1929            descending: true,
1930            nulls_first: false,
1931        };
1932
1933        let asc_ordering = LexOrdering::new(vec![PhysicalSortExpr::new(
1934            Arc::clone(&ord_expr),
1935            asc_opts,
1936        )])
1937        .unwrap();
1938        let desc_ordering = LexOrdering::new(vec![PhysicalSortExpr::new(
1939            Arc::clone(&ord_expr),
1940            desc_opts,
1941        )])
1942        .unwrap();
1943
1944        let ordering_dtype = DataType::Int64;
1945
1946        // Partial acc A: sees rows [0,1,2] arriving in ASC order (pre-ordered).
1947        // is_input_pre_ordered=true, reverse=true, ordering_req=[ASC].
1948        let mut partial_a = OrderSensitiveArrayAggAccumulator::try_new(
1949            &DataType::Int64,
1950            std::slice::from_ref(&ordering_dtype),
1951            asc_ordering.clone(),
1952            /*is_input_pre_ordered=*/ true,
1953            /*reverse=*/ true,
1954            /*ignore_nulls=*/ false,
1955        )?;
1956        let vals_a = Arc::new(Int64Array::from(vec![0i64, 1, 2])) as ArrayRef;
1957        let ords_a = Arc::new(Int64Array::from(vec![0i64, 1, 2])) as ArrayRef;
1958        partial_a.update_batch(&[vals_a, ords_a])?;
1959        let state_a = partial_a
1960            .state()?
1961            .iter()
1962            .map(|v| v.to_array())
1963            .collect::<Result<Vec<_>>>()?;
1964
1965        // Partial acc B: sees rows [3,4,5] arriving in ASC order.
1966        let mut partial_b = OrderSensitiveArrayAggAccumulator::try_new(
1967            &DataType::Int64,
1968            std::slice::from_ref(&ordering_dtype),
1969            asc_ordering,
1970            /*is_input_pre_ordered=*/ true,
1971            /*reverse=*/ true,
1972            /*ignore_nulls=*/ false,
1973        )?;
1974        let vals_b = Arc::new(Int64Array::from(vec![3i64, 4, 5])) as ArrayRef;
1975        let ords_b = Arc::new(Int64Array::from(vec![3i64, 4, 5])) as ArrayRef;
1976        partial_b.update_batch(&[vals_b, ords_b])?;
1977        let state_b = partial_b
1978            .state()?
1979            .iter()
1980            .map(|v| v.to_array())
1981            .collect::<Result<Vec<_>>>()?;
1982
1983        // Final acc: not optimized — ordering_req=[DESC], reverse=false.
1984        let mut final_acc = OrderSensitiveArrayAggAccumulator::try_new(
1985            &DataType::Int64,
1986            std::slice::from_ref(&ordering_dtype),
1987            desc_ordering,
1988            /*is_input_pre_ordered=*/ false,
1989            /*reverse=*/ false,
1990            /*ignore_nulls=*/ false,
1991        )?;
1992        final_acc.merge_batch(&state_a)?;
1993        final_acc.merge_batch(&state_b)?;
1994        let result = final_acc.evaluate()?;
1995
1996        let ScalarValue::List(list) = result else {
1997            return datafusion_common::internal_err!("expected List");
1998        };
1999        let result_vals: Vec<i64> = list
2000            .values()
2001            .as_any()
2002            .downcast_ref::<Int64Array>()
2003            .unwrap()
2004            .iter()
2005            .map(|v| v.unwrap())
2006            .collect();
2007
2008        // Expected DESC: [5, 4, 3, 2, 1, 0]
2009        assert_eq!(result_vals, vec![5i64, 4, 3, 2, 1, 0]);
2010        Ok(())
2011    }
2012
2013    struct ArrayAggAccumulatorBuilder {
2014        return_field: FieldRef,
2015        distinct: bool,
2016        order_bys: Vec<PhysicalSortExpr>,
2017        schema: Schema,
2018    }
2019
2020    impl ArrayAggAccumulatorBuilder {
2021        fn string() -> Self {
2022            Self::new(DataType::Utf8)
2023        }
2024
2025        fn new(data_type: DataType) -> Self {
2026            Self {
2027                return_field: Field::new(
2028                    "f",
2029                    DataType::List(Arc::new(Field::new_list_field(
2030                        data_type.clone(),
2031                        true,
2032                    ))),
2033                    true,
2034                )
2035                .into(),
2036                distinct: false,
2037                order_bys: vec![],
2038                schema: Schema {
2039                    fields: Fields::from(vec![Field::new("col", data_type, true)]),
2040                    metadata: Default::default(),
2041                },
2042            }
2043        }
2044
2045        fn distinct(mut self) -> Self {
2046            self.distinct = true;
2047            self
2048        }
2049
2050        fn order_by_col(mut self, col: &str, sort_options: SortOptions) -> Self {
2051            let new_order = PhysicalSortExpr::new(
2052                Arc::new(
2053                    Column::new_with_schema(col, &self.schema)
2054                        .expect("column not available in schema"),
2055                ),
2056                sort_options,
2057            );
2058            self.order_bys.push(new_order);
2059            self
2060        }
2061
2062        fn build(&self) -> Result<Box<dyn Accumulator>> {
2063            let expr = Arc::new(Column::new("col", 0));
2064            let expr_field = expr.return_field(&self.schema)?;
2065            ArrayAgg::default().accumulator(AccumulatorArgs {
2066                return_field: Arc::clone(&self.return_field),
2067                schema: &self.schema,
2068                expr_fields: &[expr_field],
2069                ignore_nulls: false,
2070                order_bys: &self.order_bys,
2071                is_reversed: false,
2072                name: "",
2073                is_distinct: self.distinct,
2074                exprs: &[expr],
2075            })
2076        }
2077
2078        fn build_two(&self) -> Result<(Box<dyn Accumulator>, Box<dyn Accumulator>)> {
2079            Ok((self.build()?, self.build()?))
2080        }
2081    }
2082
2083    fn str_arr(value: ScalarValue) -> Result<Vec<Option<String>>> {
2084        let ScalarValue::List(list) = value else {
2085            return internal_err!("ScalarValue was not a List");
2086        };
2087        Ok(as_generic_string_array::<i32>(list.values())?
2088            .iter()
2089            .map(|v| v.map(|v| v.to_string()))
2090            .collect())
2091    }
2092
2093    fn print_nulls(sort: Vec<Option<String>>) -> Vec<String> {
2094        sort.into_iter()
2095            .map(|v| v.unwrap_or_else(|| "NULL".to_string()))
2096            .collect()
2097    }
2098
2099    fn string_list_data<'a>(data: impl IntoIterator<Item = Vec<&'a str>>) -> ArrayRef {
2100        let mut builder = ListBuilder::new(StringBuilder::new());
2101        for string_list in data.into_iter() {
2102            builder.append_value(string_list.iter().map(Some).collect::<Vec<_>>());
2103        }
2104
2105        Arc::new(builder.finish())
2106    }
2107
2108    fn data<T, const N: usize>(list: [T; N]) -> ArrayRef
2109    where
2110        ScalarValue: From<T>,
2111    {
2112        let values: Vec<_> = list.into_iter().map(ScalarValue::from).collect();
2113        ScalarValue::iter_to_array(values).expect("Cannot convert to array")
2114    }
2115
2116    fn merge(
2117        mut acc1: Box<dyn Accumulator>,
2118        mut acc2: Box<dyn Accumulator>,
2119    ) -> Result<Box<dyn Accumulator>> {
2120        let intermediate_state = acc2.state().and_then(|e| {
2121            e.iter()
2122                .map(|v| v.to_array())
2123                .collect::<Result<Vec<ArrayRef>>>()
2124        })?;
2125        acc1.merge_batch(&intermediate_state)?;
2126        Ok(acc1)
2127    }
2128
2129    // ---- GroupsAccumulator tests ----
2130
2131    use arrow::array::Int32Array;
2132
2133    fn list_array_to_i32_vecs(list: &ListArray) -> Vec<Option<Vec<Option<i32>>>> {
2134        (0..list.len())
2135            .map(|i| {
2136                if list.is_null(i) {
2137                    None
2138                } else {
2139                    let arr = list.value(i);
2140                    let vals: Vec<Option<i32>> = arr
2141                        .as_any()
2142                        .downcast_ref::<Int32Array>()
2143                        .unwrap()
2144                        .iter()
2145                        .collect();
2146                    Some(vals)
2147                }
2148            })
2149            .collect()
2150    }
2151
2152    fn eval_i32_lists(
2153        acc: &mut ArrayAggGroupsAccumulator,
2154        emit_to: EmitTo,
2155    ) -> Result<Vec<Option<Vec<Option<i32>>>>> {
2156        let result = acc.evaluate(emit_to)?;
2157        Ok(list_array_to_i32_vecs(result.as_list::<i32>()))
2158    }
2159
2160    #[test]
2161    fn groups_accumulator_multiple_batches() -> Result<()> {
2162        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2163
2164        // First batch
2165        let values: ArrayRef = Arc::new(Int32Array::from(vec![1, 2, 3]));
2166        acc.update_batch(&[values], &[0, 1, 0], None, 2)?;
2167
2168        // Second batch
2169        let values: ArrayRef = Arc::new(Int32Array::from(vec![4, 5]));
2170        acc.update_batch(&[values], &[1, 0], None, 2)?;
2171
2172        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2173        assert_eq!(vals[0], Some(vec![Some(1), Some(3), Some(5)]));
2174        assert_eq!(vals[1], Some(vec![Some(2), Some(4)]));
2175
2176        Ok(())
2177    }
2178
2179    #[test]
2180    fn groups_accumulator_emit_first() -> Result<()> {
2181        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2182
2183        let values: ArrayRef = Arc::new(Int32Array::from(vec![10, 20, 30]));
2184        acc.update_batch(&[values], &[0, 1, 2], None, 3)?;
2185
2186        // Emit first 2 groups
2187        let vals = eval_i32_lists(&mut acc, EmitTo::First(2))?;
2188        assert_eq!(vals.len(), 2);
2189        assert_eq!(vals[0], Some(vec![Some(10)]));
2190        assert_eq!(vals[1], Some(vec![Some(20)]));
2191
2192        // Remaining group (was index 2, now shifted to 0)
2193        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2194        assert_eq!(vals.len(), 1);
2195        assert_eq!(vals[0], Some(vec![Some(30)]));
2196
2197        Ok(())
2198    }
2199
2200    #[test]
2201    fn groups_accumulator_emit_first_frees_batches() -> Result<()> {
2202        // Batch 0 has rows only for group 0; batch 1 has rows for
2203        // both groups. After emitting group 0, batch 0 should be
2204        // dropped entirely and batch 1 should be compacted to the
2205        // retained row(s).
2206        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2207
2208        let batch0: ArrayRef = Arc::new(Int32Array::from(vec![10, 20]));
2209        acc.update_batch(&[batch0], &[0, 0], None, 2)?;
2210
2211        let batch1: ArrayRef = Arc::new(Int32Array::from(vec![30, 40]));
2212        acc.update_batch(&[batch1], &[0, 1], None, 2)?;
2213
2214        assert_eq!(acc.batches.len(), 2);
2215        assert!(!acc.batches[0].is_empty());
2216        assert!(!acc.batches[1].is_empty());
2217
2218        // Emit group 0. Batch 0 is only referenced by group 0, so it
2219        // should be removed. Batch 1 is mixed, so it should be compacted
2220        // to contain only the retained row for group 1.
2221        let vals = eval_i32_lists(&mut acc, EmitTo::First(1))?;
2222        assert_eq!(vals[0], Some(vec![Some(10), Some(20), Some(30)]));
2223
2224        assert_eq!(acc.batches.len(), 1);
2225        let retained = acc.batches[0]
2226            .as_any()
2227            .downcast_ref::<Int32Array>()
2228            .unwrap();
2229        assert_eq!(retained.values(), &[40]);
2230        assert_eq!(acc.batch_entries, vec![vec![(0, 0)]]);
2231
2232        // Emit remaining group 1
2233        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2234        assert_eq!(vals[0], Some(vec![Some(40)]));
2235
2236        assert!(acc.batches.is_empty());
2237        assert_eq!(acc.size(), 0);
2238
2239        Ok(())
2240    }
2241
2242    #[test]
2243    fn groups_accumulator_emit_first_compacts_mixed_batches() -> Result<()> {
2244        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2245
2246        let batch: ArrayRef = Arc::new(Int32Array::from(vec![10, 20, 30, 40]));
2247        acc.update_batch(&[batch], &[0, 1, 0, 1], None, 2)?;
2248
2249        let size_before = acc.size();
2250        let vals = eval_i32_lists(&mut acc, EmitTo::First(1))?;
2251        assert_eq!(vals[0], Some(vec![Some(10), Some(30)]));
2252
2253        assert_eq!(acc.num_groups, 1);
2254        assert_eq!(acc.batches.len(), 1);
2255        let retained = acc.batches[0]
2256            .as_any()
2257            .downcast_ref::<Int32Array>()
2258            .unwrap();
2259        assert_eq!(retained.values(), &[20, 40]);
2260        assert_eq!(acc.batch_entries, vec![vec![(0, 0), (0, 1)]]);
2261        assert!(acc.size() < size_before);
2262
2263        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2264        assert_eq!(vals[0], Some(vec![Some(20), Some(40)]));
2265        assert_eq!(acc.size(), 0);
2266
2267        Ok(())
2268    }
2269
2270    #[test]
2271    fn groups_accumulator_emit_all_releases_capacity() -> Result<()> {
2272        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2273
2274        let batch: ArrayRef = Arc::new(Int32Array::from_iter_values(0..64));
2275        acc.update_batch(
2276            &[batch],
2277            &(0..64).map(|i| i % 4).collect::<Vec<_>>(),
2278            None,
2279            4,
2280        )?;
2281
2282        assert!(acc.size() > 0);
2283        let _ = eval_i32_lists(&mut acc, EmitTo::All)?;
2284
2285        assert_eq!(acc.size(), 0);
2286        assert_eq!(acc.batches.capacity(), 0);
2287        assert_eq!(acc.batch_entries.capacity(), 0);
2288
2289        Ok(())
2290    }
2291
2292    #[test]
2293    fn groups_accumulator_null_groups() -> Result<()> {
2294        // Groups that never receive values should produce null
2295        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2296
2297        let values: ArrayRef = Arc::new(Int32Array::from(vec![1]));
2298        // Only group 0 gets a value, groups 1 and 2 are empty
2299        acc.update_batch(&[values], &[0], None, 3)?;
2300
2301        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2302        assert_eq!(vals, vec![Some(vec![Some(1)]), None, None]);
2303
2304        Ok(())
2305    }
2306
2307    #[test]
2308    fn groups_accumulator_ignore_nulls() -> Result<()> {
2309        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, true);
2310
2311        let values: ArrayRef =
2312            Arc::new(Int32Array::from(vec![Some(1), None, Some(3), None]));
2313        acc.update_batch(&[values], &[0, 0, 1, 1], None, 2)?;
2314
2315        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2316        // Group 0: only non-null value is 1
2317        assert_eq!(vals[0], Some(vec![Some(1)]));
2318        // Group 1: only non-null value is 3
2319        assert_eq!(vals[1], Some(vec![Some(3)]));
2320
2321        Ok(())
2322    }
2323
2324    #[test]
2325    fn groups_accumulator_opt_filter() -> Result<()> {
2326        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2327
2328        let values: ArrayRef = Arc::new(Int32Array::from(vec![1, 2, 3, 4]));
2329        // Use a mix of false and null to filter out rows — both should
2330        // be skipped.
2331        let filter = BooleanArray::from(vec![Some(true), None, Some(true), Some(false)]);
2332        acc.update_batch(&[values], &[0, 0, 1, 1], Some(&filter), 2)?;
2333
2334        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2335        assert_eq!(vals[0], Some(vec![Some(1)])); // row 1 filtered (null)
2336        assert_eq!(vals[1], Some(vec![Some(3)])); // row 3 filtered (false)
2337
2338        Ok(())
2339    }
2340
2341    #[test]
2342    fn groups_accumulator_state_merge_roundtrip() -> Result<()> {
2343        // Accumulator 1: update_batch, then merge, then update_batch again.
2344        // Verifies that values appear in chronological insertion order.
2345        let mut acc1 = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2346        let values: ArrayRef = Arc::new(Int32Array::from(vec![1, 2]));
2347        acc1.update_batch(&[values], &[0, 1], None, 2)?;
2348
2349        // Accumulator 2
2350        let mut acc2 = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2351        let values: ArrayRef = Arc::new(Int32Array::from(vec![3, 4]));
2352        acc2.update_batch(&[values], &[0, 1], None, 2)?;
2353
2354        // Merge acc2's state into acc1
2355        let state = acc2.state(EmitTo::All)?;
2356        acc1.merge_batch(&state, &[0, 1], 2)?;
2357
2358        // Another update_batch on acc1 after the merge
2359        let values: ArrayRef = Arc::new(Int32Array::from(vec![5, 6]));
2360        acc1.update_batch(&[values], &[0, 1], None, 2)?;
2361
2362        // Each group's values in insertion order:
2363        // group 0: update(1), merge(3), update(5) → [1, 3, 5]
2364        // group 1: update(2), merge(4), update(6) → [2, 4, 6]
2365        let vals = eval_i32_lists(&mut acc1, EmitTo::All)?;
2366        assert_eq!(vals[0], Some(vec![Some(1), Some(3), Some(5)]));
2367        assert_eq!(vals[1], Some(vec![Some(2), Some(4), Some(6)]));
2368
2369        Ok(())
2370    }
2371
2372    #[test]
2373    fn groups_accumulator_convert_to_state() -> Result<()> {
2374        let acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2375
2376        let values: ArrayRef = Arc::new(Int32Array::from(vec![Some(10), None, Some(30)]));
2377        let state = acc.convert_to_state(&[values], None)?;
2378
2379        assert_eq!(state.len(), 1);
2380        let vals = list_array_to_i32_vecs(state[0].as_list::<i32>());
2381        assert_eq!(
2382            vals,
2383            vec![
2384                Some(vec![Some(10)]),
2385                Some(vec![None]), // null preserved inside list, not promoted
2386                Some(vec![Some(30)]),
2387            ]
2388        );
2389
2390        Ok(())
2391    }
2392
2393    #[test]
2394    fn groups_accumulator_convert_to_state_with_filter() -> Result<()> {
2395        let acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2396
2397        let values: ArrayRef = Arc::new(Int32Array::from(vec![10, 20, 30]));
2398        let filter = BooleanArray::from(vec![true, false, true]);
2399        let state = acc.convert_to_state(&[values], Some(&filter))?;
2400
2401        let vals = list_array_to_i32_vecs(state[0].as_list::<i32>());
2402        assert_eq!(
2403            vals,
2404            vec![
2405                Some(vec![Some(10)]),
2406                None, // filtered
2407                Some(vec![Some(30)]),
2408            ]
2409        );
2410
2411        Ok(())
2412    }
2413
2414    #[test]
2415    fn groups_accumulator_convert_to_state_merge_preserves_nulls() -> Result<()> {
2416        // Verifies that null values survive the convert_to_state -> merge_batch
2417        // round-trip when ignore_nulls is false (default null handling).
2418        let acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2419
2420        let values: ArrayRef = Arc::new(Int32Array::from(vec![Some(1), None, Some(3)]));
2421        let state = acc.convert_to_state(&[values], None)?;
2422
2423        // Feed state into a new accumulator via merge_batch
2424        let mut acc2 = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2425        acc2.merge_batch(&state, &[0, 0, 1], 2)?;
2426
2427        // Group 0 received rows 0 ([1]) and 1 ([NULL]) → [1, NULL]
2428        let vals = eval_i32_lists(&mut acc2, EmitTo::All)?;
2429        assert_eq!(vals[0], Some(vec![Some(1), None]));
2430        // Group 1 received row 2 ([3]) → [3]
2431        assert_eq!(vals[1], Some(vec![Some(3)]));
2432
2433        Ok(())
2434    }
2435
2436    #[test]
2437    fn groups_accumulator_convert_to_state_merge_ignore_nulls() -> Result<()> {
2438        // Verifies that null values are dropped in the convert_to_state ->
2439        // merge_batch round-trip when ignore_nulls is true.
2440        let acc = ArrayAggGroupsAccumulator::new(DataType::Int32, true);
2441
2442        let values: ArrayRef =
2443            Arc::new(Int32Array::from(vec![Some(1), None, Some(3), None]));
2444        let state = acc.convert_to_state(&[values], None)?;
2445
2446        let list = state[0].as_list::<i32>();
2447        // Rows 0 and 2 are valid lists; rows 1 and 3 are null list entries
2448        assert!(!list.is_null(0));
2449        assert!(list.is_null(1));
2450        assert!(!list.is_null(2));
2451        assert!(list.is_null(3));
2452
2453        // Feed state into a new accumulator via merge_batch
2454        let mut acc2 = ArrayAggGroupsAccumulator::new(DataType::Int32, true);
2455        acc2.merge_batch(&state, &[0, 0, 1, 1], 2)?;
2456
2457        // Group 0: received [1] and null (skipped) → [1]
2458        let vals = eval_i32_lists(&mut acc2, EmitTo::All)?;
2459        assert_eq!(vals[0], Some(vec![Some(1)]));
2460        // Group 1: received [3] and null (skipped) → [3]
2461        assert_eq!(vals[1], Some(vec![Some(3)]));
2462
2463        Ok(())
2464    }
2465
2466    #[test]
2467    fn groups_accumulator_all_groups_empty() -> Result<()> {
2468        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, false);
2469
2470        // Create groups but don't add any values (all filtered out)
2471        let values: ArrayRef = Arc::new(Int32Array::from(vec![1, 2]));
2472        let filter = BooleanArray::from(vec![false, false]);
2473        acc.update_batch(&[values], &[0, 1], Some(&filter), 2)?;
2474
2475        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2476        assert_eq!(vals, vec![None, None]);
2477
2478        Ok(())
2479    }
2480
2481    #[test]
2482    fn groups_accumulator_ignore_nulls_all_null_group() -> Result<()> {
2483        // When ignore_nulls is true and a group receives only nulls,
2484        // it should produce a null output
2485        let mut acc = ArrayAggGroupsAccumulator::new(DataType::Int32, true);
2486
2487        let values: ArrayRef = Arc::new(Int32Array::from(vec![None, Some(1), None]));
2488        acc.update_batch(&[values], &[0, 1, 0], None, 2)?;
2489
2490        let vals = eval_i32_lists(&mut acc, EmitTo::All)?;
2491        assert_eq!(vals[0], None); // group 0 got only nulls, all filtered
2492        assert_eq!(vals[1], Some(vec![Some(1)])); // group 1 got value 1
2493
2494        Ok(())
2495    }
2496
2497    // ---- retract_batch tests ----
2498
2499    #[test]
2500    fn retract_basic_sliding_window() -> Result<()> {
2501        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2502
2503        // Simulate ROWS BETWEEN 1 PRECEDING AND CURRENT ROW over [A, B, C, D]
2504        // Row 1: frame = [A]
2505        acc.update_batch(&[data(["A"])])?;
2506        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A"]);
2507
2508        // Row 2: frame = [A, B]
2509        acc.update_batch(&[data(["B"])])?;
2510        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2511
2512        // Row 3: frame = [B, C] — A leaves
2513        acc.update_batch(&[data(["C"])])?;
2514        acc.retract_batch(&[data(["A"])])?;
2515        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["B", "C"]);
2516
2517        // Row 4: frame = [C, D] — B leaves
2518        acc.update_batch(&[data(["D"])])?;
2519        acc.retract_batch(&[data(["B"])])?;
2520        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["C", "D"]);
2521
2522        Ok(())
2523    }
2524
2525    #[test]
2526    fn retract_multi_element_across_arrays() -> Result<()> {
2527        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2528
2529        // First batch: 3 elements
2530        acc.update_batch(&[data(["A", "B", "C"])])?;
2531        // Second batch: 1 element
2532        acc.update_batch(&[data(["D"])])?;
2533
2534        assert_eq!(
2535            print_nulls(str_arr(acc.evaluate()?)?),
2536            vec!["A", "B", "C", "D"]
2537        );
2538
2539        // Partial retract from front array: A leaves
2540        acc.retract_batch(&[data(["A"])])?;
2541        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["B", "C", "D"]);
2542
2543        // Retract spanning two arrays: B, C (rest of first array) + D (second array)
2544        acc.retract_batch(&[data(["B", "C", "D"])])?;
2545        let result = acc.evaluate()?;
2546        assert!(
2547            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2548            "expected null list after full retract, got {result:?}"
2549        );
2550
2551        Ok(())
2552    }
2553
2554    #[test]
2555    fn retract_with_nulls_preserved() -> Result<()> {
2556        // ignore_nulls = false: NULLs are stored and counted for retract
2557        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2558
2559        acc.update_batch(&[data([Some("A"), None, Some("C")])])?;
2560        assert_eq!(
2561            print_nulls(str_arr(acc.evaluate()?)?),
2562            vec!["A", "NULL", "C"]
2563        );
2564
2565        // Retract 2 elements: A and NULL both leave
2566        acc.retract_batch(&[data([Some("A"), None])])?;
2567        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["C"]);
2568
2569        Ok(())
2570    }
2571
2572    #[test]
2573    fn retract_with_ignore_nulls() -> Result<()> {
2574        // ignore_nulls = true: NULLs are NOT stored by update_batch,
2575        // so retract must only count non-null values
2576        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, true)?;
2577
2578        // update_batch with [A, NULL, C] → stores only [A, C] (NULL filtered)
2579        acc.update_batch(&[data([Some("A"), None, Some("C")])])?;
2580        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "C"]);
2581
2582        // retract_batch receives the original values including NULL: [A, NULL]
2583        // But only 1 non-null value (A) should be retracted
2584        acc.retract_batch(&[data([Some("A"), None])])?;
2585        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["C"]);
2586
2587        // retract_batch with [NULL, C] — only C (1 non-null) retracted
2588        acc.retract_batch(&[data([None, Some("C")])])?;
2589        let result = acc.evaluate()?;
2590        assert!(
2591            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2592            "expected null list after full retract, got {result:?}"
2593        );
2594
2595        Ok(())
2596    }
2597
2598    #[test]
2599    fn retract_ignore_nulls_all_nulls_batch() -> Result<()> {
2600        // When ignore_nulls = true and retract batch is all NULLs, nothing is retracted
2601        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, true)?;
2602
2603        acc.update_batch(&[data([Some("A"), Some("B")])])?;
2604        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2605
2606        // Retract batch of all NULLs: to_retract = 0, nothing changes
2607        acc.retract_batch(&[data::<Option<&str>, 3>([None, None, None])])?;
2608        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2609
2610        Ok(())
2611    }
2612
2613    #[test]
2614    fn retract_empty_accumulator() -> Result<()> {
2615        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2616
2617        // Retract on empty accumulator should be a no-op
2618        acc.retract_batch(&[data(["A"])])?;
2619        let result = acc.evaluate()?;
2620        assert!(
2621            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2622            "expected null list for empty accumulator, got {result:?}"
2623        );
2624
2625        Ok(())
2626    }
2627
2628    #[test]
2629    fn retract_front_offset_partial_consume() -> Result<()> {
2630        // Reproduces the RANGE BETWEEN 2 PRECEDING AND 2 FOLLOWING scenario:
2631        //   ts: 1, 2, 3, 4, 100
2632        //
2633        // Row 1 (ts=1): update [A,B,C] (3 elements, ts in [-1,3])
2634        // Row 2 (ts=2): update [D]     (ts=4 enters)
2635        // Row 3 (ts=3): no change      (same frame [0..4))
2636        // Row 4 (ts=4): retract [A]    (ts=1 leaves, partial consume)
2637        // Row 5 (ts=100): retract [B,C,D] (3-element retract spanning arrays)
2638        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2639
2640        // Row 1: update_batch(["A","B","C"])
2641        acc.update_batch(&[data(["A", "B", "C"])])?;
2642        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B", "C"]);
2643
2644        // Row 2: update_batch(["D"])
2645        acc.update_batch(&[data(["D"])])?;
2646        assert_eq!(
2647            print_nulls(str_arr(acc.evaluate()?)?),
2648            vec!["A", "B", "C", "D"]
2649        );
2650
2651        // Row 4: retract_batch(["A"]) — partial consume, front_offset = 1
2652        acc.retract_batch(&[data(["A"])])?;
2653        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["B", "C", "D"]);
2654
2655        // Row 5: update_batch(["E"]), then retract_batch(["B","C","D"])
2656        // retract spans: ["A","B","C"] (offset=1, 2 remaining) + ["D"] (1 element)
2657        acc.update_batch(&[data(["E"])])?;
2658        acc.retract_batch(&[data(["B", "C", "D"])])?;
2659        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["E"]);
2660
2661        Ok(())
2662    }
2663
2664    #[test]
2665    fn retract_update_after_full_drain() -> Result<()> {
2666        // Verify accumulator works correctly after being fully drained
2667        let mut acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2668
2669        acc.update_batch(&[data(["A", "B"])])?;
2670        acc.retract_batch(&[data(["A", "B"])])?;
2671
2672        // Accumulator is empty now
2673        let result = acc.evaluate()?;
2674        assert!(
2675            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2676            "expected null list, got {result:?}"
2677        );
2678
2679        // New values should work normally after drain
2680        acc.update_batch(&[data(["X", "Y"])])?;
2681        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["X", "Y"]);
2682
2683        acc.retract_batch(&[data(["X"])])?;
2684        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["Y"]);
2685
2686        Ok(())
2687    }
2688
2689    #[test]
2690    fn retract_supports_retract_batch() -> Result<()> {
2691        let acc = ArrayAggAccumulator::try_new(&DataType::Utf8, false)?;
2692        assert!(acc.supports_retract_batch());
2693
2694        let acc_ignore = ArrayAggAccumulator::try_new(&DataType::Utf8, true)?;
2695        assert!(acc_ignore.supports_retract_batch());
2696
2697        Ok(())
2698    }
2699
2700    #[test]
2701    fn retract_ignore_nulls_logical_vs_physical() -> Result<()> {
2702        // Regression test: DictionaryArray where logical nulls differ from physical nulls.
2703        // Manually construct a DictionaryArray where all indices are valid
2704        // (physical null_count = 0) but some point to null dictionary values
2705        // (logical_null_count > 0).
2706        use arrow::array::{DictionaryArray, Int32Array, StringArray};
2707
2708        let dict_type =
2709            DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8));
2710        let mut acc = ArrayAggAccumulator::try_new(&dict_type, true)?;
2711
2712        // Dictionary values: ["hello", NULL, "world"]
2713        // Keys: [0, 1, 2, 1] — all valid, but keys 1 and 3 point to null value
2714        let values = StringArray::from(vec![Some("hello"), None, Some("world")]);
2715        let keys = Int32Array::from(vec![0, 1, 2, 1]);
2716        let dict_array: ArrayRef = Arc::new(DictionaryArray::new(keys, Arc::new(values)));
2717
2718        // Confirm the divergence this test exists to exercise
2719        assert_eq!(
2720            dict_array.null_count(),
2721            0,
2722            "physical nulls: none in keys bitmap"
2723        );
2724        assert_eq!(
2725            dict_array.logical_null_count(),
2726            2,
2727            "logical nulls: keys pointing to null values"
2728        );
2729
2730        // update_batch uses logical_nulls() → stores only ["hello", "world"]
2731        acc.update_batch(std::slice::from_ref(&dict_array))?;
2732
2733        // Verify 2 elements stored
2734        let result = acc.evaluate()?;
2735        match &result {
2736            ScalarValue::List(arr) => {
2737                let values = arr.value(0);
2738                assert_eq!(values.len(), 2);
2739            }
2740            other => panic!("expected List, got {other:?}"),
2741        }
2742
2743        // retract_batch with same array: should retract 2 (logical non-nulls), not 4 (len) or 0 (physical non-nulls would be len-0=4)
2744        acc.retract_batch(&[dict_array])?;
2745        let result = acc.evaluate()?;
2746        assert!(
2747            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2748            "expected null list after full retract, got {result:?}"
2749        );
2750
2751        Ok(())
2752    }
2753
2754    #[test]
2755    fn retract_ignore_nulls_dict_partial() -> Result<()> {
2756        // Partial retraction with DictionaryArray where logical != physical nulls.
2757        // Manually construct so keys are all valid but some point to null values.
2758        use arrow::array::{DictionaryArray, Int32Array, StringArray};
2759
2760        let dict_type =
2761            DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8));
2762        let mut acc = ArrayAggAccumulator::try_new(&dict_type, true)?;
2763
2764        // update with ["A", "B", "C"] (no nulls)
2765        let values = StringArray::from(vec!["A", "B", "C"]);
2766        let keys = Int32Array::from(vec![0, 1, 2]);
2767        let update_array: ArrayRef =
2768            Arc::new(DictionaryArray::new(keys, Arc::new(values)));
2769        acc.update_batch(&[update_array])?;
2770
2771        // retract with dict ["A", NULL, NULL]:
2772        //   keys [0, 1, 1] all valid → physical null_count = 0
2773        //   keys 1,2 point to null value → logical_null_count = 2
2774        //   non-null count = 3 - 2 = 1 → retract 1 element
2775        let values = StringArray::from(vec![Some("A"), None]);
2776        let keys = Int32Array::from(vec![0, 1, 1]);
2777        let retract_array: ArrayRef =
2778            Arc::new(DictionaryArray::new(keys, Arc::new(values)));
2779
2780        assert_eq!(
2781            retract_array.null_count(),
2782            0,
2783            "physical nulls: none in keys bitmap"
2784        );
2785        assert_eq!(
2786            retract_array.logical_null_count(),
2787            2,
2788            "logical nulls: keys pointing to null values"
2789        );
2790
2791        acc.retract_batch(&[retract_array])?;
2792
2793        // Should have retracted only 1 element, leaving ["B", "C"]
2794        let result = acc.evaluate()?;
2795        match &result {
2796            ScalarValue::List(arr) => {
2797                let values = arr.value(0);
2798                assert_eq!(values.len(), 2);
2799            }
2800            other => panic!("expected List with 2 elements, got {other:?}"),
2801        }
2802
2803        Ok(())
2804    }
2805
2806    // ---- DistinctArrayAggAccumulator retract_batch tests ----
2807
2808    // Build a DISTINCT accumulator with ascending sort so evaluate output is
2809    // deterministic regardless of HashMap iteration order.
2810    fn distinct_acc(ignore_nulls: bool) -> Result<DistinctArrayAggAccumulator> {
2811        DistinctArrayAggAccumulator::try_new(
2812            &DataType::Utf8,
2813            Some(SortOptions::default()),
2814            ignore_nulls,
2815        )
2816    }
2817
2818    #[test]
2819    fn distinct_retract_duplicate_remains() -> Result<()> {
2820        // Canonical regression for the HashSet-can't-retract bug: a value
2821        // that appears multiple times in-frame must survive retraction of
2822        // a single occurrence.
2823        let mut acc = distinct_acc(false)?;
2824
2825        // Feed [A, A, B] across two batches to exercise multi-batch state.
2826        acc.update_batch(&[data(["A", "A"])])?;
2827        acc.update_batch(&[data(["B"])])?;
2828        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2829
2830        // Retract a single A — the other A is still in the frame.
2831        acc.retract_batch(&[data(["A"])])?;
2832        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2833
2834        // Retract the remaining A — only B left.
2835        acc.retract_batch(&[data(["A"])])?;
2836        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["B"]);
2837
2838        Ok(())
2839    }
2840
2841    #[test]
2842    fn distinct_retract_full_removal() -> Result<()> {
2843        let mut acc = distinct_acc(false)?;
2844
2845        acc.update_batch(&[data(["A", "B"])])?;
2846        acc.retract_batch(&[data(["A", "B"])])?;
2847
2848        let result = acc.evaluate()?;
2849        assert!(
2850            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
2851            "expected null list after full retract, got {result:?}"
2852        );
2853
2854        Ok(())
2855    }
2856
2857    #[test]
2858    fn distinct_retract_ignore_nulls_skips() -> Result<()> {
2859        // ignore_nulls=true: NULL never enters state on update, so retract
2860        // must also skip NULL — otherwise we'd error on the missing key.
2861        let mut acc = distinct_acc(true)?;
2862
2863        acc.update_batch(&[data([Some("A"), None, Some("B")])])?;
2864        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A", "B"]);
2865
2866        // Retract [A, NULL] — the NULL is skipped, only A is removed.
2867        acc.retract_batch(&[data([Some("A"), None])])?;
2868        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["B"]);
2869
2870        Ok(())
2871    }
2872
2873    #[test]
2874    fn distinct_retract_null_tracked() -> Result<()> {
2875        // ignore_nulls=false: NULL enters state with a refcount and must
2876        // retract symmetrically; the NULL key must be removed at zero
2877        // (else evaluate still emits a NULL element).
2878        let mut acc = distinct_acc(false)?;
2879
2880        acc.update_batch(&[data([Some("A"), None, None])])?;
2881        // With nulls_first=true (SortOptions default), NULL sorts before A.
2882        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["NULL", "A"]);
2883
2884        // Retract one NULL — count drops to 1, key still present.
2885        acc.retract_batch(&[data::<Option<&str>, 1>([None])])?;
2886        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["NULL", "A"]);
2887
2888        // Retract the remaining NULL — key is removed.
2889        acc.retract_batch(&[data::<Option<&str>, 1>([None])])?;
2890        assert_eq!(print_nulls(str_arr(acc.evaluate()?)?), vec!["A"]);
2891
2892        Ok(())
2893    }
2894
2895    #[test]
2896    fn distinct_supports_retract_batch() -> Result<()> {
2897        let acc = distinct_acc(false)?;
2898        assert!(acc.supports_retract_batch());
2899
2900        let acc_ignore = distinct_acc(true)?;
2901        assert!(acc_ignore.supports_retract_batch());
2902
2903        Ok(())
2904    }
2905
2906    #[test]
2907    fn distinct_merge_then_evaluate_regression() -> Result<()> {
2908        // Non-window path: state -> merge_batch -> evaluate must still
2909        // produce the union of distinct values across partitions.
2910        let mut acc1 = distinct_acc(false)?;
2911        let mut acc2 = distinct_acc(false)?;
2912
2913        acc1.update_batch(&[data(["A", "A", "B"])])?;
2914        acc2.update_batch(&[data(["A", "C"])])?;
2915
2916        let state = acc2.state()?;
2917        let state_arrs: Vec<ArrayRef> = state
2918            .into_iter()
2919            .map(|sv| sv.to_array_of_size(1))
2920            .collect::<Result<Vec<_>>>()?;
2921        acc1.merge_batch(&state_arrs)?;
2922
2923        assert_eq!(print_nulls(str_arr(acc1.evaluate()?)?), vec!["A", "B", "C"]);
2924
2925        Ok(())
2926    }
2927
2928    #[test]
2929    fn distinct_array_agg_utf8_deduplicates() -> Result<()> {
2930        use arrow::array::StringArray;
2931
2932        // 7 rows with 4 distinct values, each duplicate appearing twice.
2933        let input: ArrayRef = Arc::new(StringArray::from(vec![
2934            "postgres", "mysql", "postgres", "redis", "mysql", "duckdb", "redis",
2935        ]));
2936
2937        let mut acc = DistinctArrayAggAccumulator::try_new(&DataType::Utf8, None, false)?;
2938        acc.update_batch(&[input])?;
2939
2940        let result = acc.evaluate()?;
2941        let ScalarValue::List(arr) = &result else {
2942            panic!("expected ScalarValue::List, got {result:?}");
2943        };
2944
2945        let inner = arr.value(0);
2946        let strings = inner
2947            .as_any()
2948            .downcast_ref::<StringArray>()
2949            .expect("inner array should be StringArray");
2950
2951        // HashSet ordering is nondeterministic — sort before asserting.
2952        let mut values: Vec<&str> =
2953            (0..strings.len()).map(|i| strings.value(i)).collect();
2954        values.sort_unstable();
2955
2956        assert_eq!(values, vec!["duckdb", "mysql", "postgres", "redis"]);
2957        Ok(())
2958    }
2959
2960    #[test]
2961    fn distinct_array_agg_int64_deduplicates() -> Result<()> {
2962        use arrow::array::Int64Array;
2963
2964        // 7 rows with 4 distinct values, each duplicate appearing twice.
2965        let input: ArrayRef = Arc::new(Int64Array::from(vec![1i64, 2, 1, 3, 2, 4, 3]));
2966
2967        let mut acc =
2968            DistinctArrayAggAccumulator::try_new(&DataType::Int64, None, false)?;
2969        acc.update_batch(&[input])?;
2970
2971        let result = acc.evaluate()?;
2972        let ScalarValue::List(arr) = &result else {
2973            panic!("expected ScalarValue::List, got {result:?}");
2974        };
2975
2976        let inner = arr.value(0);
2977        let ints = inner
2978            .as_any()
2979            .downcast_ref::<Int64Array>()
2980            .expect("inner array should be Int64Array");
2981
2982        let mut values: Vec<i64> = (0..ints.len()).map(|i| ints.value(i)).collect();
2983        values.sort_unstable();
2984
2985        assert_eq!(values, vec![1i64, 2, 3, 4]);
2986        Ok(())
2987    }
2988
2989    #[test]
2990    fn distinct_array_agg_float64_deduplicates() -> Result<()> {
2991        use arrow::array::Float64Array;
2992
2993        // 7 rows with 4 distinct values, each duplicate appearing twice.
2994        let input: ArrayRef = Arc::new(Float64Array::from(vec![
2995            1.0f64, 2.5, 1.0, 3.75, 2.5, 4.0, 3.75,
2996        ]));
2997
2998        let mut acc =
2999            DistinctArrayAggAccumulator::try_new(&DataType::Float64, None, false)?;
3000        acc.update_batch(&[input])?;
3001
3002        let result = acc.evaluate()?;
3003        let ScalarValue::List(arr) = &result else {
3004            panic!("expected ScalarValue::List, got {result:?}");
3005        };
3006
3007        let inner = arr.value(0);
3008        let floats = inner
3009            .as_any()
3010            .downcast_ref::<Float64Array>()
3011            .expect("inner array should be Float64Array");
3012
3013        // f64 has no Ord — use total_cmp for a stable sort.
3014        let mut values: Vec<f64> = (0..floats.len()).map(|i| floats.value(i)).collect();
3015        values.sort_unstable_by(|a, b| a.total_cmp(b));
3016
3017        assert_eq!(values, vec![1.0f64, 2.5, 3.75, 4.0]);
3018        Ok(())
3019    }
3020
3021    #[test]
3022    fn distinct_array_agg_dictionary_preserves_type() -> Result<()> {
3023        use arrow::array::{DictionaryArray, Int32Array, StringArray};
3024
3025        // Dictionary(Int32, Utf8) input with duplicates.
3026        let keys = Int32Array::from(vec![0, 1, 0, 2, 1]); // "a", "b", "a", "c", "b"
3027        let values = StringArray::from(vec!["a", "b", "c"]);
3028        let dict: ArrayRef = Arc::new(DictionaryArray::new(keys, Arc::new(values)));
3029
3030        let datatype =
3031            DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8));
3032        let mut acc = DistinctArrayAggAccumulator::try_new(&datatype, None, false)?;
3033        acc.update_batch(&[dict])?;
3034
3035        let result = acc.evaluate()?;
3036        let ScalarValue::List(arr) = &result else {
3037            panic!("expected ScalarValue::List, got {result:?}");
3038        };
3039
3040        // The element type of the returned list must stay Dictionary(Int32, Utf8),
3041        // not be silently widened to Utf8.
3042        assert_eq!(
3043            arr.values().data_type(),
3044            &datatype,
3045            "element type must be Dictionary(Int32, Utf8), got {}",
3046            arr.values().data_type()
3047        );
3048
3049        // There should be exactly 3 distinct values.
3050        assert_eq!(arr.value(0).len(), 3);
3051        Ok(())
3052    }
3053
3054    #[test]
3055    fn distinct_array_agg_date32_deduplicates() -> Result<()> {
3056        use arrow::array::Date32Array;
3057
3058        // 7 rows with 4 distinct dates (days since epoch), each duplicate appearing twice.
3059        let input: ArrayRef = Arc::new(Date32Array::from(vec![
3060            100i32, 200, 100, 300, 200, 400, 300,
3061        ]));
3062
3063        let mut acc =
3064            DistinctArrayAggAccumulator::try_new(&DataType::Date32, None, false)?;
3065        acc.update_batch(&[input])?;
3066
3067        let result = acc.evaluate()?;
3068        let ScalarValue::List(arr) = &result else {
3069            panic!("expected ScalarValue::List, got {result:?}");
3070        };
3071
3072        let inner = arr.value(0);
3073        let dates = inner
3074            .as_any()
3075            .downcast_ref::<Date32Array>()
3076            .expect("inner array should be Date32Array");
3077
3078        let mut values: Vec<i32> = (0..dates.len()).map(|i| dates.value(i)).collect();
3079        values.sort_unstable();
3080
3081        assert_eq!(values, vec![100i32, 200, 300, 400]);
3082        Ok(())
3083    }
3084
3085    #[test]
3086    fn distinct_retract_memory_is_bounded() -> Result<()> {
3087        use arrow::array::Int64Array;
3088
3089        // Emulates a sliding window where each value enters and immediately
3090        // leaves. Only CARDINALITY distinct values are ever live at once;
3091        // memory must not grow with the number of rows processed.
3092        const CARDINALITY: i64 = 10;
3093        const WARMUP_ROWS: i64 = 1_000;
3094        const EXTRA_ROWS: i64 = 20_000;
3095
3096        let mut acc =
3097            DistinctArrayAggAccumulator::try_new(&DataType::Int64, None, false)?;
3098
3099        let slide = |acc: &mut DistinctArrayAggAccumulator, rows: i64| -> Result<()> {
3100            for i in 0..rows {
3101                let value: ArrayRef = Arc::new(Int64Array::from(vec![i % CARDINALITY]));
3102                acc.update_batch(std::slice::from_ref(&value))?;
3103                acc.retract_batch(std::slice::from_ref(&value))?;
3104            }
3105            Ok(())
3106        };
3107
3108        // Let every buffer reach its steady state before taking a baseline.
3109        slide(&mut acc, WARMUP_ROWS)?;
3110        let baseline = acc.size();
3111
3112        slide(&mut acc, EXTRA_ROWS)?;
3113        let grown = acc.size();
3114
3115        assert!(
3116            grown <= 2 * baseline,
3117            "size() must not grow with the number of retracted rows: \
3118             {baseline} bytes after {WARMUP_ROWS} rows, \
3119             {grown} bytes after {} rows",
3120            WARMUP_ROWS + EXTRA_ROWS
3121        );
3122
3123        // Everything was retracted so evaluate must return null.
3124        let result = acc.evaluate()?;
3125        assert!(
3126            matches!(&result, ScalarValue::List(arr) if arr.is_null(0)),
3127            "expected null list after retracting every row, got {result:?}"
3128        );
3129
3130        Ok(())
3131    }
3132}