kumarUjjawal commented on code in PR #25740: URL: https://github.com/apache/datafusion/pull/25740#discussion_r4176419459
########## datafusion/functions-aggregate/src/map_agg.rs: ########## @@ -0,0 +1,991 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! `MAP_AGG` aggregate implementation: [`MapAgg`] +//! +//! Aggregates key/value pairs into a `Map`, analogous to how `array_agg` +//! aggregates values into a `List`. +//! +//! # Accumulators +//! +//! Two accumulators implement the function, selected in +//! [`MapAgg::accumulator`] by whether the aggregate carries an `ORDER BY`: +//! +//! * [`MapAggAccumulator`] keeps the pairs in input order and serializes its +//! state as a single `Map` column. +//! * [`OrderSensitiveMapAggAccumulator`] additionally records the values of the +//! ordering expressions for every pair. Its state therefore has two columns, +//! the `Map` plus a `List<Struct<ordering...>>` that lets partial states from +//! every partition be merged on the ordering by `merge_batch`, so the final +//! entry order stays deterministic under parallelism. +//! +//! [`MapAgg::order_sensitivity`] is [`AggregateOrderSensitivity::SoftRequirement`], +//! so when a plan's input already satisfies the ordering the accumulator is +//! marked through [`AggregateUDFImpl::with_beneficial_ordering`] and the pairs +//! are not sorted again. +//! +//! # De-duplication +//! +//! Arrow maps hold at most one value per key, so duplicate keys are removed +//! when the accumulator evaluates. The *first* pair wins, where "first" means +//! first in input order for [`MapAggAccumulator`] and first in the sorted +//! order for [`OrderSensitiveMapAggAccumulator`]. +//! +//! # Nulls +//! +//! Arrow maps cannot represent `NULL` keys: [`MapAggAccumulator`] fails with +//! `map key cannot be null` when one is evaluated, while +//! [`OrderSensitiveMapAggAccumulator`] drops such pairs as it collects input. +//! `NULL` values are ordinary map values and are kept. An empty group +//! evaluates to a `NULL` map, matching `array_agg`. + +use std::collections::{HashSet, VecDeque}; +use std::mem::{size_of, size_of_val, take}; +use std::sync::Arc; + +use arrow::array::{Array, ArrayRef, AsArray, MapArray, StructArray}; +use arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer}; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, FieldRef, Fields}; + +use datafusion_common::cast::as_map_array; +use datafusion_common::utils::{ + SingleRowListArrayBuilder, compare_rows, get_row_at_idx, take_function_args, +}; +use datafusion_common::{ + Result, ScalarValue, assert_eq_or_internal_err, exec_err, internal_err, +}; +use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs}; +use datafusion_expr::utils::format_state_name; +use datafusion_expr::{ + Accumulator, AggregateUDFImpl, Documentation, Signature, Volatility, +}; +use datafusion_functions_aggregate_common::merge_arrays::merge_ordered_arrays; +use datafusion_functions_aggregate_common::order::AggregateOrderSensitivity; +use datafusion_functions_aggregate_common::utils::ordering_fields; +use datafusion_macros::user_doc; +use datafusion_physical_expr_common::sort_expr::LexOrdering; + +use crate::utils::{map_row_to_scalars, struct_to_rows}; + +make_udaf_expr_and_func!( + MapAgg, + map_agg, + key value, + "Aggregates keys and values into a map", + map_agg_udaf +); + +#[user_doc( + doc_section(label = "General Functions"), + description = "Returns a map created from the key and value expression elements. \ +For each row, the key expression becomes a map key and the value expression becomes the corresponding map value. \ +Entries appear in input order, or in the order given by the optional `ORDER BY`. \ +When a key repeats, only the first entry for that key is kept.", + syntax_example = "map_agg(key, value [ORDER BY expression])", + sql_example = r#"```sql +> SELECT map_agg(column_key, column_value) FROM table_name; ++-------------------------------------+ +| map_agg(column_key, column_value) | ++-------------------------------------+ +| {key1: value1, key2: value2, ...} | ++-------------------------------------+ +```"#, + argument( + name = "key", + description = "Expression used as the map key. Can be a column or any valid expression." + ), + argument( + name = "value", + description = "Expression used as the map value. Can be a column or any valid expression." + ) +)] +#[derive(Debug, PartialEq, Eq, Hash)] +/// MAP_AGG aggregate expression +pub struct MapAgg { + /// Accepts two arguments of any type: the key and the value. + signature: Signature, + /// Whether the input is known to arrive already ordered by the `ORDER BY` + /// inside the aggregate. + is_input_pre_ordered: bool, +} + +impl Default for MapAgg { + fn default() -> Self { + Self { + signature: Signature::any(2, Volatility::Immutable), + is_input_pre_ordered: false, + } + } +} + +impl MapAgg { + /// Create a new MAP_AGG aggregate function + pub fn new() -> Self { + Self::default() + } + + /// Build the Arrow `Map` data type for `map_agg(key, value)` + fn map_data_type(key_type: &DataType, value_type: &DataType) -> DataType { + DataType::Map( + Arc::new(Field::new( + "entries", + DataType::Struct(Fields::from(vec![ + Field::new("key", key_type.clone(), false), + Field::new("value", value_type.clone(), true), + ])), + false, + )), + false, + ) + } +} + +impl AggregateUDFImpl for MapAgg { + fn name(&self) -> &str { + "map_agg" + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> { + let [key_type, value_type] = take_function_args(self.name(), arg_types)?; + Ok(Self::map_data_type(key_type, value_type)) + } + + fn state_fields(&self, args: StateFieldsArgs) -> Result<Vec<FieldRef>> { + let map_type = Self::map_data_type( + args.input_fields[0].data_type(), + args.input_fields[1].data_type(), + ); + + let mut fields = vec![ + Field::new( + format_state_name(args.name, "map_agg"), + // Nullable so empty groups can produce a NULL map + map_type, + true, + ) + .into(), + ]; + + if args.ordering_fields.is_empty() { + return Ok(fields); + } + + let orderings = args.ordering_fields.to_vec(); + fields.push( + Field::new_list( + format_state_name(args.name, "map_agg_orderings"), + Field::new_list_field(DataType::Struct(Fields::from(orderings)), true), + false, + ) + .into(), + ); + + Ok(fields) + } + + fn accumulator(&self, acc_args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> { + let [key_field, value_field] = + take_function_args(self.name(), acc_args.expr_fields)?; + let key_type = key_field.data_type().clone(); + let value_type = value_field.data_type().clone(); + + let Some(ordering) = LexOrdering::new(acc_args.order_bys.to_vec()) else { + return MapAggAccumulator::try_new(key_type, value_type) + .map(|acc| Box::new(acc) as _); + }; + + let ordering_dtypes = ordering + .iter() + .map(|e| e.expr.data_type(acc_args.schema)) + .collect::<Result<Vec<_>>>()?; + + Ok(Box::new(OrderSensitiveMapAggAccumulator::new( + key_type, + value_type, + ordering_dtypes, + ordering, + self.is_input_pre_ordered, + ))) + } + + fn order_sensitivity(&self) -> AggregateOrderSensitivity { + AggregateOrderSensitivity::SoftRequirement + } + + fn with_beneficial_ordering( + self: Arc<Self>, + beneficial_ordering: bool, + ) -> Result<Option<Arc<dyn AggregateUDFImpl>>> { + Ok(Some(Arc::new(Self { + signature: self.signature.clone(), + is_input_pre_ordered: beneficial_ordering, + }))) + } + + fn documentation(&self) -> Option<&Documentation> { + self.doc() + } +} + +/// Accumulates key/value pairs for [`MapAgg`]. +/// +/// Input rows are kept as parallel [`ScalarValue`] lists and concatenated into +/// the map in [`Self::evaluate`]; partial states are single-row [`MapArray`] +/// scalars that are split back into rows in [`Self::merge_batch`]. +#[derive(Debug)] +pub struct MapAggAccumulator { + /// Data type of the map keys. + key_type: DataType, + /// Data type of the map values. + value_type: DataType, + /// Input keys in arrival order, parallel to `values`. + keys: Vec<ScalarValue>, + /// Input values in arrival order, parallel to `keys`. + values: Vec<ScalarValue>, +} + +impl MapAggAccumulator { + /// Create a new map_agg accumulator for the given key and value types + pub fn try_new(key_type: DataType, value_type: DataType) -> Result<Self> { + Ok(Self { + key_type, + value_type, + keys: vec![], + values: vec![], + }) + } + /// Keeps only the first occurrence of each key, preserving input order. + /// + /// Returns the surviving keys and values as two aligned vectors. + #[allow(clippy::allow_attributes, clippy::mutable_key_type)] // ScalarValue has interior mutability but is intentionally used as hash key + fn dedup_first_wins( + keys: Vec<ScalarValue>, + values: Vec<ScalarValue>, + ) -> (Vec<ScalarValue>, Vec<ScalarValue>) { + // First pass: mark each position that is the first occurrence of its key. + let mut seen = HashSet::with_capacity(keys.len()); + let keep: Vec<bool> = keys.iter().map(|k| seen.insert(k.clone())).collect(); + + // Second pass: keep only the first-occurrence positions. + let out_keys = keys + .into_iter() + .zip(&keep) + .filter_map(|(k, &keep)| keep.then_some(k)) + .collect(); + let out_values = values + .into_iter() + .zip(&keep) + .filter_map(|(v, &keep)| keep.then_some(v)) + .collect(); + (out_keys, out_values) + } + /// The `DataType::Map` this accumulator evaluates to + fn map_type(&self) -> DataType { + MapAgg::map_data_type(&self.key_type, &self.value_type) + } +} + +impl Accumulator for MapAggAccumulator { + fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> { + assert_eq_or_internal_err!(values.len(), 2, "map_agg expects (key, value)"); + + let keys = &values[0]; + let vals = &values[1]; + assert_eq_or_internal_err!(keys.len(), vals.len(), "key/value length mismatch"); + + for row in 0..keys.len() { + self.keys.push(ScalarValue::try_from_array(keys, row)?); + self.values.push(ScalarValue::try_from_array(vals, row)?); Review Comment: Compact nested scalars before storing them, as the ordered accumulator's update path already does, and apply the same treatment when decoding partial states. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
