Skip to main content

common_function/aggrs/
aggr_wrapper.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Wrapper for making aggregate functions out of state/merge functions of original aggregate functions.
16//!
17//! i.e. for a aggregate function `foo`, we will have a state function `foo_state` and a merge function `foo_merge`.
18//!
19//! `foo_state`'s input args is the same as `foo`'s, and its output is a state object.
20//! Note that `foo_state` might have multiple output columns, so it's a struct array
21//! that each output column is a struct field.
22//! `foo_merge`'s input arg is the same as `foo_state`'s output, and its output is the same as `foo`'s input.
23//!
24
25use std::hash::{Hash, Hasher};
26use std::sync::Arc;
27
28use arrow::array::{ArrayData, ArrayRef, BooleanArray, StructArray, make_array};
29use arrow_schema::{FieldRef, Fields};
30use common_telemetry::debug;
31use datafusion::functions_aggregate::all_default_aggregate_functions;
32use datafusion::functions_aggregate::count::Count;
33use datafusion::functions_aggregate::min_max::{Max, Min};
34use datafusion::optimizer::AnalyzerRule;
35use datafusion::optimizer::analyzer::type_coercion::TypeCoercion;
36use datafusion_common::{Column, ScalarValue};
37use datafusion_expr::expr::{AggregateFunction, AggregateFunctionParams};
38use datafusion_expr::function::StateFieldsArgs;
39use datafusion_expr::physical_planning_context::PhysicalPlanningContext;
40use datafusion_expr::{
41    Accumulator, Aggregate, AggregateUDF, AggregateUDFImpl, EmitTo, Expr, ExprSchemable,
42    GroupsAccumulator, LogicalPlan, Signature,
43};
44use datafusion_physical_expr::aggregate::{AggregateFunctionExpr, LoweredAggregateBuilder};
45use datatypes::arrow::datatypes::{DataType, Field};
46
47use crate::aggrs::aggr_wrapper::fix_order::FixStateUdafOrderingAnalyzer;
48use crate::function_registry::{FUNCTION_REGISTRY, FunctionRegistry};
49
50pub mod fix_order;
51#[cfg(test)]
52mod tests;
53
54/// Returns the name of the state function for the given aggregate function name.
55/// The state function is used to compute the state of the aggregate function.
56/// The state function's name is in the format `__<aggr_name>_state
57pub fn aggr_state_func_name(aggr_name: &str) -> String {
58    format!("__{}_state", aggr_name)
59}
60
61/// Returns the name of the merge function for the given aggregate function name.
62/// The merge function is used to merge the states of the state functions.
63/// The merge function's name is in the format `__<aggr_name>_merge
64pub fn aggr_merge_func_name(aggr_name: &str) -> String {
65    format!("__{}_merge", aggr_name)
66}
67
68/// Returns the globally registered name used to merge a delta state with a
69/// persisted state.
70pub fn aggr_delta_merge_func_name(state_aggregate_name: &str) -> String {
71    format!("__{}_delta_merge", state_aggregate_name)
72}
73
74/// Check if the given aggregate expression is steppable.
75/// As in if it can be split into multiple steps:
76/// i.e. on datanode first call `state(input)` then
77/// on frontend call `calc(merge(state))` to get the final result.
78pub fn is_all_aggr_exprs_steppable(aggr_exprs: &[Expr]) -> bool {
79    aggr_exprs.iter().all(|expr| {
80        if let Some(aggr_func) = get_aggr_func(expr) {
81            if aggr_func.params.distinct {
82                // Distinct aggregate functions are not steppable(yet).
83                // TODO(discord9): support distinct aggregate functions.
84                return false;
85            }
86
87            // DataFusion only sorts the input of an aggregate with a hard ordering requirement
88            // when the requirement is already satisfied or the aggregate has a reverse
89            // expression (apache/datafusion#25676). The state wrapper has none, so e.g.
90            // `nth_value(.. ORDER BY ..)` would read unsorted input on datanodes. Ordered-set
91            // aggregates like `approx_percentile_cont(..) WITHIN GROUP (ORDER BY ..)` are
92            // exempt: their ORDER BY names the value, and they don't need sorted input.
93            if !aggr_func.params.order_by.is_empty()
94                && aggr_func.func.order_sensitivity().hard_requires()
95                && !aggr_func.func.supports_within_group_clause()
96            {
97                return false;
98            }
99
100            // whether the corresponding state function exists in the registry
101            FUNCTION_REGISTRY.is_aggr_func_exist(&aggr_state_func_name(aggr_func.func.name()))
102        } else {
103            false
104        }
105    })
106}
107
108pub fn get_aggr_func(expr: &Expr) -> Option<&datafusion_expr::expr::AggregateFunction> {
109    let mut expr_ref = expr;
110    while let Expr::Alias(alias) = expr_ref {
111        expr_ref = &alias.expr;
112    }
113    if let Expr::AggregateFunction(aggr_func) = expr_ref {
114        Some(aggr_func)
115    } else {
116        None
117    }
118}
119
120/// A wrapper to make an aggregate function out of the state and merge functions of the original aggregate function.
121/// It contains the original aggregate function, the state functions, and the merge function.
122///
123/// Notice state functions may have multiple output columns, so it's return type is always a struct array, and the merge function is used to merge the states of the state functions.
124#[derive(Debug, Clone)]
125pub struct StateMergeHelper;
126
127/// A struct to hold the two aggregate plans, one for the state function(lower) and one for the merge function(upper).
128#[allow(unused)]
129#[derive(Debug, Clone)]
130pub struct StepAggrPlan {
131    /// Upper merge plan, which is the aggregate plan that merges the states of the state function.
132    pub upper_merge: LogicalPlan,
133    /// Lower state plan, which is the aggregate plan that computes the state of the aggregate function.
134    pub lower_state: LogicalPlan,
135}
136
137impl StateMergeHelper {
138    /// Register all the `state` function of supported aggregate functions.
139    /// Note that can't register `merge` function here, as it needs to be created from the original aggregate function with given input types.
140    pub fn register(registry: &FunctionRegistry) {
141        let all_default = all_default_aggregate_functions();
142        let greptime_custom_aggr_functions = registry.aggregate_functions();
143
144        // if our custom aggregate function have the same name as the default aggregate function, we will override it.
145        let supported = all_default
146            .into_iter()
147            .chain(greptime_custom_aggr_functions.into_iter().map(Arc::new))
148            .collect::<Vec<_>>();
149        debug!(
150            "Registering state functions for supported: {:?}",
151            supported.iter().map(|f| f.name()).collect::<Vec<_>>()
152        );
153
154        let state_func = supported.into_iter().filter_map(|f| {
155            StateWrapper::new((*f).clone())
156                .inspect_err(
157                    |e| common_telemetry::error!(e; "Failed to register state function for {:?}", f),
158                )
159                .ok()
160                .map(AggregateUDF::new_from_impl)
161        });
162
163        for func in state_func {
164            registry.register_aggr(func);
165        }
166    }
167
168    /// Split an aggregate plan into two aggregate plans, one for the state function and one for the merge function.
169    ///
170    pub fn split_aggr_node(aggr_plan: Aggregate) -> datafusion_common::Result<StepAggrPlan> {
171        let aggr = {
172            // certain aggr func need type coercion to work correctly, so we need to analyze the plan first.
173            let aggr_plan = TypeCoercion::new().analyze(
174                LogicalPlan::Aggregate(aggr_plan).clone(),
175                &Default::default(),
176            )?;
177            if let LogicalPlan::Aggregate(aggr) = aggr_plan {
178                aggr
179            } else {
180                return Err(datafusion_common::DataFusionError::Internal(format!(
181                    "Failed to coerce expressions in aggregate plan, expected Aggregate, got: {:?}",
182                    aggr_plan
183                )));
184            }
185        };
186        let mut lower_aggr_exprs = vec![];
187        let mut upper_aggr_exprs = vec![];
188
189        // group exprs for upper plan should refer to the output group expr as column from lower plan
190        // to avoid re-compute group exprs again.
191        let upper_group_exprs = aggr
192            .group_expr
193            .iter()
194            .map(|c| c.qualified_name())
195            .map(|(r, c)| Expr::Column(Column::new(r, c)))
196            .collect();
197
198        for aggr_expr in aggr.aggr_expr.iter() {
199            let Some(aggr_func) = get_aggr_func(aggr_expr) else {
200                return Err(datafusion_common::DataFusionError::NotImplemented(format!(
201                    "Unsupported aggregate expression for step aggr optimize: {:?}",
202                    aggr_expr
203                )));
204            };
205
206            let original_input_fields = aggr_func
207                .params
208                .args
209                .iter()
210                .map(|e| e.to_field(&aggr.input.schema()).map(|(_, field)| field))
211                .collect::<Result<Vec<_>, _>>()?;
212
213            // first create the state function from the original aggregate function.
214            let state_func = StateWrapper::new((*aggr_func.func).clone())?;
215
216            let expr = AggregateFunction {
217                func: Arc::new(state_func.into()),
218                params: aggr_func.params.clone(),
219            };
220            let expr = Expr::AggregateFunction(expr);
221            let lower_state_output_col_name = expr.schema_name().to_string();
222
223            lower_aggr_exprs.push(expr);
224
225            // then create the merge function using the physical expression of the original aggregate function
226            let (name, human_display) = match aggr_expr {
227                Expr::Alias(alias) => (alias.name.clone(), aggr_expr.human_display().to_string()),
228                Expr::AggregateFunction(_) => (
229                    aggr_expr.schema_name().to_string(),
230                    aggr_expr.human_display().to_string(),
231                ),
232                _ => unreachable!("aggregate expression was validated above"),
233            };
234            let original_phy_expr = LoweredAggregateBuilder::new(
235                aggr_expr,
236                aggr.input.schema(),
237                aggr.input.schema().as_arrow(),
238                &Default::default(),
239                &PhysicalPlanningContext::default(),
240            )
241            .with_name(name)
242            .with_human_display(human_display)
243            .build()?
244            .aggregate;
245
246            let merge_func = MergeWrapper::new(
247                (*aggr_func.func).clone(),
248                original_phy_expr,
249                original_input_fields,
250            )?;
251            let arg = Expr::Column(Column::new_unqualified(lower_state_output_col_name));
252            let expr = AggregateFunction {
253                func: Arc::new(merge_func.into()),
254                // notice filter/order_by is not supported in the merge function, as it's not meaningful to have them in the merge phase.
255                // do notice this order by is only removed in the outer logical plan, the physical plan still have order by and hence
256                // can create correct accumulator with order by.
257                params: AggregateFunctionParams {
258                    args: vec![arg],
259                    distinct: aggr_func.params.distinct,
260                    filter: None,
261                    order_by: vec![],
262                    null_treatment: aggr_func.params.null_treatment,
263                },
264            };
265
266            // alias to the original aggregate expr's schema name, so parent plan can refer to it
267            // correctly.
268            let expr = Expr::AggregateFunction(expr).alias(aggr_expr.schema_name().to_string());
269            upper_aggr_exprs.push(expr);
270        }
271
272        let mut lower = aggr.clone();
273        lower.aggr_expr = lower_aggr_exprs;
274        let lower_plan = LogicalPlan::Aggregate(lower);
275
276        // update aggregate's output schema
277        let lower_plan = lower_plan.recompute_schema()?;
278
279        // should only affect two udaf `first_value/last_value`
280        // which only them have meaningful order by field
281        let fixed_lower_plan =
282            FixStateUdafOrderingAnalyzer.analyze(lower_plan, &Default::default())?;
283
284        let upper = Aggregate::try_new(
285            Arc::new(fixed_lower_plan.clone()),
286            upper_group_exprs,
287            upper_aggr_exprs.clone(),
288        )?;
289        let aggr_plan = LogicalPlan::Aggregate(aggr);
290
291        // upper schema's output schema should be the same as the original aggregate plan's output schema
292        let upper_check = upper;
293        let upper_plan = LogicalPlan::Aggregate(upper_check).recompute_schema()?;
294        if *upper_plan.schema() != *aggr_plan.schema() {
295            return Err(datafusion_common::DataFusionError::Internal(format!(
296                "Upper aggregate plan's schema is not the same as the original aggregate plan's schema: \n[transformed]:{}\n[original]:{}",
297                upper_plan.schema(),
298                aggr_plan.schema()
299            )));
300        }
301
302        Ok(StepAggrPlan {
303            lower_state: fixed_lower_plan,
304            upper_merge: upper_plan,
305        })
306    }
307}
308
309/// Wrapper to make an aggregate function out of a state function.
310#[derive(Debug, Clone, PartialEq, Eq, Hash)]
311pub struct StateWrapper {
312    inner: AggregateUDF,
313    name: String,
314    /// Default to empty, might get fixed by analyzer later
315    ordering: Vec<FieldRef>,
316    /// Default to false, might get fixed by analyzer later
317    distinct: bool,
318}
319
320impl StateWrapper {
321    /// `state_index`: The index of the state in the output of the state function.
322    pub fn new(inner: AggregateUDF) -> datafusion_common::Result<Self> {
323        let name = aggr_state_func_name(inner.name());
324        Ok(Self {
325            inner,
326            name,
327            ordering: vec![],
328            distinct: false,
329        })
330    }
331
332    pub fn inner(&self) -> &AggregateUDF {
333        &self.inner
334    }
335
336    /// Deduce the return type of the original aggregate function
337    /// based on the accumulator arguments.
338    ///
339    pub fn deduce_aggr_return_type(
340        &self,
341        acc_args: &datafusion_expr::function::AccumulatorArgs,
342    ) -> datafusion_common::Result<FieldRef> {
343        let input_fields = acc_args
344            .exprs
345            .iter()
346            .map(|e| e.return_field(acc_args.schema))
347            .collect::<Result<Vec<_>, _>>()?;
348        self.inner.return_field(&input_fields).inspect_err(|e| {
349            common_telemetry::error!(
350                "StateWrapper: {:#?}\nacc_args:{:?}\nerror:{:?}",
351                &self,
352                &acc_args,
353                e
354            );
355        })
356    }
357
358    fn fix_inner_acc_args<'b>(
359        &self,
360        mut acc_args: datafusion_expr::function::AccumulatorArgs<'b>,
361    ) -> datafusion_common::Result<datafusion_expr::function::AccumulatorArgs<'b>> {
362        acc_args.return_field = self.deduce_aggr_return_type(&acc_args)?;
363        Ok(acc_args)
364    }
365}
366
367impl AggregateUDFImpl for StateWrapper {
368    fn accumulator<'a, 'b>(
369        &'a self,
370        acc_args: datafusion_expr::function::AccumulatorArgs<'b>,
371    ) -> datafusion_common::Result<Box<dyn Accumulator>> {
372        // fix and recover proper acc args for the original aggregate function.
373        let state_type = acc_args.return_type().clone();
374        let inner = self.inner.accumulator(self.fix_inner_acc_args(acc_args)?)?;
375
376        Ok(Box::new(StateAccum::new(inner, state_type)?))
377    }
378
379    fn groups_accumulator_supported(
380        &self,
381        acc_args: datafusion_expr::function::AccumulatorArgs,
382    ) -> bool {
383        self.fix_inner_acc_args(acc_args)
384            .map(|args| self.inner.inner().groups_accumulator_supported(args))
385            .unwrap_or(false)
386    }
387
388    fn create_groups_accumulator(
389        &self,
390        acc_args: datafusion_expr::function::AccumulatorArgs,
391    ) -> datafusion_common::Result<Box<dyn GroupsAccumulator>> {
392        let state_type = acc_args.return_type().clone();
393        let inner = self
394            .inner
395            .inner()
396            .create_groups_accumulator(self.fix_inner_acc_args(acc_args)?)?;
397        Ok(Box::new(StateGroupsAccum::new(inner, state_type)?))
398    }
399
400    fn name(&self) -> &str {
401        self.name.as_str()
402    }
403
404    fn is_nullable(&self) -> bool {
405        self.inner.is_nullable()
406    }
407
408    /// Return state_fields as the output struct type.
409    ///
410    fn return_type(&self, arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
411        let input_fields = &arg_types
412            .iter()
413            .map(|x| Arc::new(Field::new("x", x.clone(), false)))
414            .collect::<Vec<_>>();
415
416        let state_fields_args = StateFieldsArgs {
417            name: self.inner().name(),
418            input_fields,
419            return_field: self.inner.return_field(input_fields)?,
420            // those args are also needed as they are vital to construct the state fields correctly.
421            ordering_fields: &self.ordering,
422            is_distinct: self.distinct,
423        };
424        let state_fields = self.inner.state_fields(state_fields_args)?;
425
426        let state_fields = state_fields
427            .into_iter()
428            .map(|f| {
429                let mut f = f.as_ref().clone();
430                // since state can be null when no input rows, so make all fields nullable
431                f.set_nullable(true);
432                Arc::new(f)
433            })
434            .collect::<Vec<_>>();
435
436        let struct_field = DataType::Struct(state_fields.into());
437        Ok(struct_field)
438    }
439
440    /// The state function's output fields are the same as the original aggregate function's state fields.
441    fn state_fields(
442        &self,
443        args: datafusion_expr::function::StateFieldsArgs,
444    ) -> datafusion_common::Result<Vec<FieldRef>> {
445        let state_fields_args = StateFieldsArgs {
446            name: args.name,
447            input_fields: args.input_fields,
448            return_field: self.inner.return_field(args.input_fields)?,
449            ordering_fields: args.ordering_fields,
450            is_distinct: args.is_distinct,
451        };
452        self.inner.state_fields(state_fields_args)
453    }
454
455    /// The state function's signature is the same as the original aggregate function's signature,
456    fn signature(&self) -> &Signature {
457        self.inner.signature()
458    }
459
460    /// Coerce types also do nothing, as optimizer should be able to already make struct types
461    fn coerce_types(&self, arg_types: &[DataType]) -> datafusion_common::Result<Vec<DataType>> {
462        self.inner.coerce_types(arg_types)
463    }
464
465    fn value_from_stats(
466        &self,
467        statistics_args: &datafusion_expr::StatisticsArgs,
468    ) -> Option<ScalarValue> {
469        let inner = self.inner().inner();
470        // only count/min/max need special handling here, for getting result from statistics
471        // the result of count/min/max is also the result of count_state so can return directly
472        let can_use_stat = inner.is::<Count>() || inner.is::<Max>() || inner.is::<Min>();
473        if !can_use_stat {
474            return None;
475        }
476
477        // fix return type by extract the first field's data type from the struct type
478        let state_type = if let DataType::Struct(fields) = &statistics_args.return_type {
479            if fields.is_empty() {
480                return None;
481            }
482            fields[0].data_type().clone()
483        } else {
484            return None;
485        };
486
487        let fixed_args = datafusion_expr::StatisticsArgs {
488            statistics: statistics_args.statistics,
489            return_type: &state_type,
490            is_distinct: statistics_args.is_distinct,
491            exprs: statistics_args.exprs,
492        };
493
494        let ret = self.inner().value_from_stats(&fixed_args)?;
495
496        // wrap the result into struct scalar value
497        let fields = if let DataType::Struct(fields) = &statistics_args.return_type {
498            fields
499        } else {
500            return None;
501        };
502
503        let array = ret.to_array().ok()?;
504
505        let struct_array = StructArray::new(fields.clone(), vec![array], None);
506        let ret = ScalarValue::Struct(Arc::new(struct_array));
507        Some(ret)
508    }
509}
510
511/// The wrapper's input is the same as the original aggregate function's input,
512/// and the output is the state function's output.
513#[derive(Debug)]
514pub struct StateAccum {
515    inner: Box<dyn Accumulator>,
516    state_fields: Fields,
517}
518
519pub struct StateGroupsAccum {
520    inner: Box<dyn GroupsAccumulator>,
521    state_fields: Fields,
522}
523
524impl StateGroupsAccum {
525    fn new(
526        inner: Box<dyn GroupsAccumulator>,
527        state_type: DataType,
528    ) -> datafusion_common::Result<Self> {
529        let DataType::Struct(fields) = state_type else {
530            return Err(datafusion_common::DataFusionError::Internal(format!(
531                "Expected a struct type for state, got: {:?}",
532                state_type
533            )));
534        };
535        Ok(Self {
536            inner,
537            state_fields: fields,
538        })
539    }
540
541    fn wrap_state_arrays(&self, arrays: Vec<ArrayRef>) -> datafusion_common::Result<ArrayRef> {
542        Ok(Arc::new(state_struct_array(&self.state_fields, arrays)?))
543    }
544}
545
546impl GroupsAccumulator for StateGroupsAccum {
547    fn update_batch(
548        &mut self,
549        values: &[ArrayRef],
550        group_indices: &[usize],
551        opt_filter: Option<&BooleanArray>,
552        total_num_groups: usize,
553    ) -> datafusion_common::Result<()> {
554        self.inner
555            .update_batch(values, group_indices, opt_filter, total_num_groups)
556    }
557
558    fn merge_batch(
559        &mut self,
560        values: &[ArrayRef],
561        group_indices: &[usize],
562        total_num_groups: usize,
563    ) -> datafusion_common::Result<()> {
564        self.inner
565            .merge_batch(values, group_indices, total_num_groups)
566    }
567
568    fn evaluate(&mut self, emit_to: EmitTo) -> datafusion_common::Result<ArrayRef> {
569        let state = self.inner.state(emit_to)?;
570        self.wrap_state_arrays(state)
571    }
572
573    fn state(&mut self, emit_to: EmitTo) -> datafusion_common::Result<Vec<ArrayRef>> {
574        self.inner.state(emit_to)
575    }
576
577    fn convert_to_state(
578        &self,
579        values: &[ArrayRef],
580        opt_filter: Option<&BooleanArray>,
581    ) -> datafusion_common::Result<Vec<ArrayRef>> {
582        self.inner.convert_to_state(values, opt_filter)
583    }
584
585    fn size(&self) -> usize {
586        self.inner.size()
587    }
588}
589
590/// Wraps the state arrays of an accumulator into a struct of the declared state fields.
591///
592/// The declared state type is derived from logical expressions, while the accumulator names
593/// nested fields after physical expressions. For example, `array_agg(v ORDER BY ts)` declares
594/// its orderings as `List(Struct("ts": ..))` but produces `List(Struct("ts@0": ..))`. Arrays that
595/// differ only in nested field names are cast to the declared type; any other difference is an
596/// error.
597fn state_struct_array(
598    state_fields: &Fields,
599    arrays: Vec<ArrayRef>,
600) -> datafusion_common::Result<StructArray> {
601    if arrays.len() != state_fields.len() {
602        return Err(datafusion_common::DataFusionError::Internal(format!(
603            "Expected {} state arrays for fields {:?}, got {}",
604            state_fields.len(),
605            state_fields,
606            arrays.len()
607        )));
608    }
609    let arrays = arrays
610        .into_iter()
611        .zip(state_fields.iter())
612        .map(|(array, field)| {
613            let expected = field.data_type();
614            if array.data_type() == expected {
615                Ok(array)
616            } else if array.data_type().equals_datatype(expected) {
617                Ok(make_array(relabel_nested_fields(
618                    array.to_data(),
619                    expected,
620                )?))
621            } else {
622                Err(datafusion_common::DataFusionError::Internal(format!(
623                    "State field `{}` expects type {expected}, but the accumulator produced {}",
624                    field.name(),
625                    array.data_type()
626                )))
627            }
628        })
629        .collect::<datafusion_common::Result<Vec<_>>>()?;
630    Ok(StructArray::try_new(state_fields.clone(), arrays, None)?)
631}
632
633/// Rebuilds `data` with the type `target`, which must match it position by position apart
634/// from nested field names and metadata. Unlike a cast, children are never matched by name.
635fn relabel_nested_fields(
636    data: ArrayData,
637    target: &DataType,
638) -> datafusion_common::Result<ArrayData> {
639    let child_types = match target {
640        DataType::List(field)
641        | DataType::LargeList(field)
642        | DataType::FixedSizeList(field, _)
643        | DataType::Map(field, _) => vec![field.data_type()],
644        DataType::Struct(fields) => fields.iter().map(|f| f.data_type()).collect(),
645        _ if data.child_data().is_empty() => vec![],
646        _ => {
647            return Err(datafusion_common::DataFusionError::NotImplemented(format!(
648                "Relabeling nested fields of {target}"
649            )));
650        }
651    };
652    let children = data
653        .child_data()
654        .iter()
655        .zip(child_types)
656        .map(|(child, child_type)| relabel_nested_fields(child.clone(), child_type))
657        .collect::<datafusion_common::Result<Vec<_>>>()?;
658    Ok(data
659        .into_builder()
660        .data_type(target.clone())
661        .child_data(children)
662        .build()?)
663}
664
665impl StateAccum {
666    pub fn new(
667        inner: Box<dyn Accumulator>,
668        state_type: DataType,
669    ) -> datafusion_common::Result<Self> {
670        let DataType::Struct(fields) = state_type else {
671            return Err(datafusion_common::DataFusionError::Internal(format!(
672                "Expected a struct type for state, got: {:?}",
673                state_type
674            )));
675        };
676        Ok(Self {
677            inner,
678            state_fields: fields,
679        })
680    }
681}
682
683impl Accumulator for StateAccum {
684    fn evaluate(&mut self) -> datafusion_common::Result<ScalarValue> {
685        let state = self.inner.state()?;
686
687        let array = state
688            .iter()
689            .map(|s| s.to_array())
690            .collect::<Result<Vec<_>, _>>()?;
691        let struct_array = state_struct_array(&self.state_fields, array)?;
692        Ok(ScalarValue::Struct(Arc::new(struct_array)))
693    }
694
695    fn merge_batch(
696        &mut self,
697        states: &[datatypes::arrow::array::ArrayRef],
698    ) -> datafusion_common::Result<()> {
699        self.inner.merge_batch(states)
700    }
701
702    fn update_batch(
703        &mut self,
704        values: &[datatypes::arrow::array::ArrayRef],
705    ) -> datafusion_common::Result<()> {
706        self.inner.update_batch(values)
707    }
708
709    fn size(&self) -> usize {
710        self.inner.size()
711    }
712
713    fn state(&mut self) -> datafusion_common::Result<Vec<ScalarValue>> {
714        self.inner.state()
715    }
716}
717
718/// A globally registerable wrapper for a state-family merge UDAF.
719///
720/// The wrapped merge function has the family contract `P..., State`. This
721/// wrapper exposes `P..., delta_state, persisted_state` and forwards the two
722/// state columns to the existing accumulator in that order. State values are
723/// opaque to this adapter; null is the inner merge family's identity value.
724#[derive(Debug, Clone)]
725pub(crate) struct DeltaMergeWrapper {
726    inner: AggregateUDF,
727    name: String,
728    signature: Signature,
729    inner_types: Vec<DataType>,
730    state_type: DataType,
731}
732
733impl DeltaMergeWrapper {
734    /// Build the wrapper for one of the explicitly supported exact merge UDAFs.
735    ///
736    /// The caller supplies the known merge signature, so construction cannot
737    /// fail while inspecting an arbitrary UDAF signature.
738    pub(crate) fn new(
739        inner: AggregateUDF,
740        state_name: &str,
741        inner_types: Vec<DataType>,
742        state_type: DataType,
743    ) -> Self {
744        let mut wrapper_types = inner_types.clone();
745        wrapper_types.push(state_type.clone());
746        Self {
747            name: aggr_delta_merge_func_name(state_name),
748            signature: Signature::exact(wrapper_types, inner.signature().volatility),
749            inner,
750            inner_types,
751            state_type,
752        }
753    }
754
755    fn resolve_inner_args(
756        &self,
757        input_fields: &[FieldRef],
758    ) -> datafusion_common::Result<Vec<FieldRef>> {
759        if input_fields.len() != self.inner_types.len() + 1 {
760            return Err(datafusion_common::DataFusionError::Plan(
761                "delta merge requires parameters, delta state, and persisted state".to_string(),
762            ));
763        }
764        for (field, expected_type) in input_fields[..self.inner_types.len()]
765            .iter()
766            .zip(&self.inner_types)
767        {
768            if field.data_type() != expected_type {
769                return Err(datafusion_common::DataFusionError::Plan(format!(
770                    "delta merge argument type does not match its exact signature: {:?} != {expected_type:?}",
771                    field.data_type()
772                )));
773            }
774        }
775        let persisted = &input_fields[self.inner_types.len()];
776        if persisted.data_type() != &self.state_type && persisted.data_type() != &DataType::Null {
777            return Err(datafusion_common::DataFusionError::Plan(format!(
778                "persisted state type does not match the exact state type: {:?} != {:?}",
779                persisted.data_type(),
780                self.state_type
781            )));
782        }
783        Ok(input_fields[..self.inner_types.len()].to_vec())
784    }
785}
786
787impl AggregateUDFImpl for DeltaMergeWrapper {
788    fn accumulator<'a, 'b>(
789        &'a self,
790        acc_args: datafusion_expr::function::AccumulatorArgs<'b>,
791    ) -> datafusion_common::Result<Box<dyn Accumulator>> {
792        if acc_args.exprs.len() != acc_args.expr_fields.len() {
793            return Err(datafusion_common::DataFusionError::Plan(
794                "delta merge expression and field arities differ".to_string(),
795            ));
796        }
797        let inner_fields = self.resolve_inner_args(acc_args.expr_fields)?;
798        for (expr, expected_type) in acc_args.exprs.iter().zip(
799            self.inner_types
800                .iter()
801                .chain(std::iter::once(&self.state_type)),
802        ) {
803            if expr.data_type(acc_args.schema)? != *expected_type {
804                return Err(datafusion_common::DataFusionError::Internal(
805                    "delta merge physical expression type is not resolved".to_string(),
806                ));
807            }
808        }
809        let state_index = self.inner_types.len() - 1;
810        let inner_args = datafusion_expr::function::AccumulatorArgs {
811            return_field: acc_args.return_field,
812            schema: acc_args.schema,
813            ignore_nulls: acc_args.ignore_nulls,
814            order_bys: acc_args.order_bys,
815            is_reversed: acc_args.is_reversed,
816            name: self.inner.name(),
817            is_distinct: acc_args.is_distinct,
818            exprs: &acc_args.exprs[..=state_index],
819            expr_fields: &inner_fields,
820        };
821        Ok(Box::new(DeltaMergeAccum {
822            inner: self.inner.accumulator(inner_args)?,
823            params: state_index,
824        }))
825    }
826
827    fn name(&self) -> &str {
828        &self.name
829    }
830
831    fn is_nullable(&self) -> bool {
832        self.inner.is_nullable()
833    }
834
835    fn return_type(&self, arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
836        let fields = arg_types
837            .iter()
838            .enumerate()
839            .map(|(index, data_type)| {
840                Arc::new(Field::new(index.to_string(), data_type.clone(), true))
841            })
842            .collect::<Vec<_>>();
843        let inner_fields = self.resolve_inner_args(&fields)?;
844        self.inner.return_type(
845            &inner_fields
846                .iter()
847                .map(|field| field.data_type().clone())
848                .collect::<Vec<_>>(),
849        )
850    }
851
852    fn return_field(&self, arg_fields: &[FieldRef]) -> datafusion_common::Result<FieldRef> {
853        let inner_fields = self.resolve_inner_args(arg_fields)?;
854        self.inner.return_field(&inner_fields)
855    }
856
857    fn signature(&self) -> &Signature {
858        &self.signature
859    }
860
861    fn state_fields(
862        &self,
863        args: datafusion_expr::function::StateFieldsArgs,
864    ) -> datafusion_common::Result<Vec<FieldRef>> {
865        let inner_fields = self.resolve_inner_args(args.input_fields)?;
866        self.inner
867            .state_fields(datafusion_expr::function::StateFieldsArgs {
868                name: args.name,
869                input_fields: &inner_fields,
870                return_field: args.return_field,
871                ordering_fields: args.ordering_fields,
872                is_distinct: args.is_distinct,
873            })
874    }
875}
876
877impl PartialEq for DeltaMergeWrapper {
878    fn eq(&self, other: &Self) -> bool {
879        self.name == other.name && self.inner == other.inner
880    }
881}
882impl Eq for DeltaMergeWrapper {}
883impl Hash for DeltaMergeWrapper {
884    fn hash<H: Hasher>(&self, state: &mut H) {
885        self.name.hash(state);
886        self.inner.hash(state);
887    }
888}
889
890#[derive(Debug)]
891struct DeltaMergeAccum {
892    inner: Box<dyn Accumulator>,
893    params: usize,
894}
895
896impl Accumulator for DeltaMergeAccum {
897    fn evaluate(&mut self) -> datafusion_common::Result<ScalarValue> {
898        self.inner.evaluate()
899    }
900
901    fn update_batch(&mut self, values: &[ArrayRef]) -> datafusion_common::Result<()> {
902        if values.len() != self.params + 2 {
903            return Err(datafusion_common::DataFusionError::Plan(format!(
904                "delta merge expected {} arguments, got {}",
905                self.params + 2,
906                values.len()
907            )));
908        }
909        // Null-state identity is an inner family precondition. The wrapper
910        // forwards opaque states and never decodes or filters them.
911        let mut inner_values = values[..self.params + 1].to_vec();
912        self.inner.update_batch(&inner_values)?;
913        inner_values[self.params] = values[self.params + 1].clone();
914        self.inner.update_batch(&inner_values)
915    }
916
917    fn merge_batch(&mut self, states: &[ArrayRef]) -> datafusion_common::Result<()> {
918        self.inner.merge_batch(states)
919    }
920
921    fn size(&self) -> usize {
922        self.inner.size()
923    }
924
925    fn state(&mut self) -> datafusion_common::Result<Vec<ScalarValue>> {
926        self.inner.state()
927    }
928}
929
930/// TODO(discord9): mark this function as non-ser/de able
931///
932/// This wrapper shouldn't be register as a udaf, as it contain extra data that is not serializable.
933/// and changes for different logical plans.
934#[derive(Debug, Clone)]
935pub struct MergeWrapper {
936    inner: AggregateUDF,
937    name: String,
938    merge_signature: Signature,
939    /// The original physical expression of the aggregate function, can't store the original aggregate function directly, as PhysicalExpr didn't implement Any
940    original_phy_expr: Arc<AggregateFunctionExpr>,
941    return_field: FieldRef,
942}
943impl MergeWrapper {
944    pub fn new(
945        inner: AggregateUDF,
946        original_phy_expr: Arc<AggregateFunctionExpr>,
947        original_input_fields: Vec<FieldRef>,
948    ) -> datafusion_common::Result<Self> {
949        let name = aggr_merge_func_name(inner.name());
950        // the input type is actually struct type, which is the state fields of the original aggregate function.
951        let merge_signature = Signature::user_defined(datafusion_expr::Volatility::Immutable);
952        let return_field = inner.return_field(&original_input_fields)?.clone();
953
954        Ok(Self {
955            inner,
956            name,
957            merge_signature,
958            original_phy_expr,
959            return_field,
960        })
961    }
962
963    pub fn inner(&self) -> &AggregateUDF {
964        &self.inner
965    }
966}
967
968impl AggregateUDFImpl for MergeWrapper {
969    fn accumulator<'a, 'b>(
970        &'a self,
971        acc_args: datafusion_expr::function::AccumulatorArgs<'b>,
972    ) -> datafusion_common::Result<Box<dyn Accumulator>> {
973        if acc_args.exprs.len() != 1
974            || !matches!(
975                acc_args.exprs[0].data_type(acc_args.schema)?,
976                DataType::Struct(_)
977            )
978        {
979            return Err(datafusion_common::DataFusionError::Internal(format!(
980                "Expected one struct type as input, got: {:?}",
981                acc_args.schema
982            )));
983        }
984        let input_type = acc_args.exprs[0].data_type(acc_args.schema)?;
985        let DataType::Struct(fields) = input_type else {
986            return Err(datafusion_common::DataFusionError::Internal(format!(
987                "Expected a struct type for input, got: {:?}",
988                input_type
989            )));
990        };
991
992        let inner_accum = self.original_phy_expr.create_accumulator()?;
993        Ok(Box::new(MergeAccum::new(inner_accum, &fields)))
994    }
995
996    fn name(&self) -> &str {
997        self.name.as_str()
998    }
999
1000    fn is_nullable(&self) -> bool {
1001        self.inner.is_nullable()
1002    }
1003
1004    /// Notice here the `arg_types` is actually the `state_fields`'s data types,
1005    /// so return fixed return type instead of using `arg_types` to determine the return type.
1006    fn return_type(&self, _arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
1007        // The return type is the same as the original aggregate function's return type.
1008        Ok(self.return_field.data_type().clone())
1009    }
1010
1011    /// Similar to return_type, we just return the fixed return field.
1012    fn return_field(&self, _arg_fields: &[FieldRef]) -> datafusion_common::Result<FieldRef> {
1013        Ok(self.return_field.clone())
1014    }
1015
1016    fn signature(&self) -> &Signature {
1017        &self.merge_signature
1018    }
1019
1020    /// Coerce types also do nothing, as optimizer should be able to already make struct types
1021    fn coerce_types(&self, arg_types: &[DataType]) -> datafusion_common::Result<Vec<DataType>> {
1022        // just check if the arg_types are only one and is struct array
1023        if arg_types.len() != 1 || !matches!(arg_types.first(), Some(DataType::Struct(_))) {
1024            return Err(datafusion_common::DataFusionError::Internal(format!(
1025                "Expected one struct type as input, got: {:?}",
1026                arg_types
1027            )));
1028        }
1029        Ok(arg_types.to_vec())
1030    }
1031
1032    /// Just return the original aggregate function's state fields.
1033    fn state_fields(
1034        &self,
1035        _args: datafusion_expr::function::StateFieldsArgs,
1036    ) -> datafusion_common::Result<Vec<FieldRef>> {
1037        self.original_phy_expr.state_fields()
1038    }
1039}
1040
1041impl PartialEq for MergeWrapper {
1042    fn eq(&self, other: &Self) -> bool {
1043        self.inner == other.inner
1044    }
1045}
1046
1047impl Eq for MergeWrapper {}
1048
1049impl Hash for MergeWrapper {
1050    fn hash<H: Hasher>(&self, state: &mut H) {
1051        self.inner.hash(state);
1052    }
1053}
1054
1055/// The merge accumulator, which modify `update_batch`'s behavior to accept one struct array which
1056/// include the state fields of original aggregate function, and merge said states into original accumulator
1057/// the output is the same as original aggregate function
1058#[derive(Debug)]
1059pub struct MergeAccum {
1060    inner: Box<dyn Accumulator>,
1061    state_fields: Fields,
1062}
1063
1064impl MergeAccum {
1065    pub fn new(inner: Box<dyn Accumulator>, state_fields: &Fields) -> Self {
1066        Self {
1067            inner,
1068            state_fields: state_fields.clone(),
1069        }
1070    }
1071}
1072
1073impl Accumulator for MergeAccum {
1074    fn evaluate(&mut self) -> datafusion_common::Result<ScalarValue> {
1075        self.inner.evaluate()
1076    }
1077
1078    fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> datafusion_common::Result<()> {
1079        self.inner.merge_batch(states)
1080    }
1081
1082    fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> datafusion_common::Result<()> {
1083        let value = values.first().ok_or_else(|| {
1084            datafusion_common::DataFusionError::Internal("No values provided for merge".to_string())
1085        })?;
1086        // The input values are states from other accumulators, so we merge them.
1087        let struct_arr = value
1088            .as_any()
1089            .downcast_ref::<StructArray>()
1090            .ok_or_else(|| {
1091                datafusion_common::DataFusionError::Internal(format!(
1092                    "Expected StructArray, got: {:?}",
1093                    value.data_type()
1094                ))
1095            })?;
1096        let fields = struct_arr.fields();
1097        if fields != &self.state_fields {
1098            debug!(
1099                "State fields mismatch, expected: {:?}, got: {:?}",
1100                self.state_fields, fields
1101            );
1102            // state fields mismatch might be acceptable by datafusion, continue
1103        }
1104
1105        // now fields should be the same, so we can merge the batch
1106        // by pass the columns as order should be the same
1107        let state_columns = struct_arr.columns();
1108        self.inner.merge_batch(state_columns)
1109    }
1110
1111    fn size(&self) -> usize {
1112        self.inner.size()
1113    }
1114
1115    fn state(&mut self) -> datafusion_common::Result<Vec<ScalarValue>> {
1116        self.inner.state()
1117    }
1118}