1use 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
54pub fn aggr_state_func_name(aggr_name: &str) -> String {
58 format!("__{}_state", aggr_name)
59}
60
61pub fn aggr_merge_func_name(aggr_name: &str) -> String {
65 format!("__{}_merge", aggr_name)
66}
67
68pub fn aggr_delta_merge_func_name(state_aggregate_name: &str) -> String {
71 format!("__{}_delta_merge", state_aggregate_name)
72}
73
74pub 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 return false;
85 }
86
87 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 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#[derive(Debug, Clone)]
125pub struct StateMergeHelper;
126
127#[allow(unused)]
129#[derive(Debug, Clone)]
130pub struct StepAggrPlan {
131 pub upper_merge: LogicalPlan,
133 pub lower_state: LogicalPlan,
135}
136
137impl StateMergeHelper {
138 pub fn register(registry: &FunctionRegistry) {
141 let all_default = all_default_aggregate_functions();
142 let greptime_custom_aggr_functions = registry.aggregate_functions();
143
144 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 pub fn split_aggr_node(aggr_plan: Aggregate) -> datafusion_common::Result<StepAggrPlan> {
171 let aggr = {
172 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 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 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 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 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 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 let lower_plan = lower_plan.recompute_schema()?;
278
279 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 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#[derive(Debug, Clone, PartialEq, Eq, Hash)]
311pub struct StateWrapper {
312 inner: AggregateUDF,
313 name: String,
314 ordering: Vec<FieldRef>,
316 distinct: bool,
318}
319
320impl StateWrapper {
321 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 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 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 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 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 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 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 fn signature(&self) -> &Signature {
457 self.inner.signature()
458 }
459
460 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 let can_use_stat = inner.is::<Count>() || inner.is::<Max>() || inner.is::<Min>();
473 if !can_use_stat {
474 return None;
475 }
476
477 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 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#[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
590fn 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
633fn 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#[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 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 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#[derive(Debug, Clone)]
935pub struct MergeWrapper {
936 inner: AggregateUDF,
937 name: String,
938 merge_signature: Signature,
939 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 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 fn return_type(&self, _arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
1007 Ok(self.return_field.data_type().clone())
1009 }
1010
1011 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 fn coerce_types(&self, arg_types: &[DataType]) -> datafusion_common::Result<Vec<DataType>> {
1022 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 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#[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 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 }
1104
1105 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}