namanjain24-sudo commented on code in PR #25527:
URL: https://github.com/apache/datafusion/pull/25527#discussion_r4210179021


##########
datafusion/substrait/src/logical_plan/grouping_set.rs:
##########
@@ -0,0 +1,204 @@
+// 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.
+
+//! The column a multi-set aggregate ends with, which Substrait and DataFusion
+//! fill differently.
+//!
+//! Substrait gives an [`AggregateRel`] with more than one grouping set a
+//! trailing `i32` holding "the zero-based index of the grouping set that
+//! yielded the record" ([Aggregate Operation]). DataFusion ends the same
+//! aggregate with `__grouping_id`, which packs a bitmask of the columns the 
set
+//! leaves out together with an ordinal separating repeated sets. Both identify
+//! the set a row came from, so each side can be written as a map of the other,
+//! which is what the consumer and the producer apply.
+//!
+//! [`AggregateRel`]: substrait::proto::AggregateRel
+//! [Aggregate Operation]: 
https://substrait.io/relations/logical_relations/#aggregate-operation
+
+use datafusion::arrow::datatypes::DataType;
+use datafusion::common::{
+    Column, ScalarValue, internal_datafusion_err, internal_err, not_impl_err,
+};
+use datafusion::logical_expr::utils::grouping_set_to_exprlist;
+use datafusion::logical_expr::{Aggregate, Case, Expr, GroupingSet, lit};
+
+/// The name the grouping set index is given, which Substrait leaves to the
+/// plan's root names.
+pub(crate) const GROUPING_SET_INDEX: &str = "grouping_set_index";
+
+/// The `__grouping_id` value DataFusion gives each grouping set, in the order
+/// the sets are listed.
+///
+/// The value is `(ordinal << group_count) | mask`: a bit is set in `mask` for
+/// every grouping column the set leaves out, counting from the last column, 
and
+/// `ordinal` counts the sets before this one holding the same columns. Both
+/// parts follow from the set alone, so no two sets share a value.
+pub(crate) fn grouping_set_ids(
+    columns: &[&Expr],
+    sets: &[Vec<Expr>],
+) -> datafusion::common::Result<Vec<u64>> {
+    let group_count = columns.len();
+    if group_count > 64 {
+        return not_impl_err!(
+            "Grouping sets with more than 64 columns are not supported"
+        );
+    }
+
+    let mut ids = Vec::with_capacity(sets.len());
+    let mut masks = Vec::with_capacity(sets.len());
+    for set in sets {
+        let mut mask = 0u64;
+        for (position, column) in columns.iter().enumerate() {
+            if !set.contains(column) {
+                mask |= 1 << (group_count - 1 - position);
+            }
+        }
+        let ordinal = masks.iter().filter(|seen| **seen == mask).count() as 
u64;
+        masks.push(mask);
+        ids.push((ordinal << group_count) | mask);

Review Comment:
   Fixed: at group_count==64 the zero-ordinal case now skips the shift, and a 
repeated set at 64 columns returns a clean error instead of panicking. Added 
producer/consumer regressions plus unit tests on `grouping_set_ids` directly.
   
   Also found while testing: DataFusion's own native aggregate execution 
(`physical-plan/src/aggregates/mod.rs`, around line 3216) has the exact same 
unconditional `ordinal << n` at `n == 64`, independent of Substrait. Out of 
scope here, but flagging in case it's worth its own issue.



##########
datafusion/substrait/src/logical_plan/producer/rel/aggregate_rel.rs:
##########
@@ -39,35 +47,93 @@ pub fn from_aggregate(
         .iter()
         .map(|e| to_substrait_agg_measure(producer, e, agg.input.schema()))
         .collect::<datafusion::common::Result<Vec<_>>>()?;
-    let common = (groupings.len() > 1)
-        .then(|| grouping_set_output_mapping(grouping_expressions.len(), 
measures.len()));
-
-    Ok(Box::new(Rel {
+    let is_grouping_set = groupings.len() > 1;
+    let grouping_count = grouping_expressions.len();
+    let measure_count = measures.len();
+    let aggregate = Box::new(Rel {
         rel_type: Some(RelType::Aggregate(Box::new(AggregateRel {
-            common,
+            common: None,
             input: Some(input),
             grouping_expressions,
             groupings,
             measures,
             advanced_extension: None,
         }))),
-    }))
+    });
+
+    if is_grouping_set {
+        grouping_set_projection(producer, agg, aggregate, grouping_count, 
measure_count)
+    } else {
+        Ok(aggregate)
+    }
 }
 
-/// Maps Substrait's `[groups, measures, grouping_id]` direct output to
-/// DataFusion's `[groups, grouping_id, measures]` aggregate schema.
-fn grouping_set_output_mapping(grouping_count: usize, measure_count: usize) -> 
RelCommon {
-    let grouping_id_index = grouping_count + measure_count;
+/// Puts a multi-set aggregate back into DataFusion's
+/// `[groups, grouping_id, measures]` schema.
+///
+/// Substrait's own output is `[groups, measures, grouping set index]`, so
+/// besides the reordering the trailing column has to be mapped back to
+/// `__grouping_id`; see [`crate::logical_plan::grouping_set`]. That map is a
+/// projected expression, which leaves the `AggregateRel` itself holding the
+/// index the spec defines, for a consumer that reads it.
+fn grouping_set_projection(
+    producer: &mut impl SubstraitProducer,
+    agg: &Aggregate,
+    aggregate: Box<Rel>,
+    grouping_count: usize,
+    measure_count: usize,
+) -> datafusion::common::Result<Box<Rel>> {
+    let schema = agg.schema.as_ref();
+    let (grouping_id_index, _) = grouping_id_column(schema)?;
+    if grouping_id_index != grouping_count {
+        return internal_err!(
+            "Aggregate has {grouping_id_index} grouping columns but 
{grouping_count} grouping expressions were written"
+        );
+    }
+    let grouping_id_type = schema.field(grouping_id_index).data_type();
+    let ids = grouping_set_ids(
+        &grouping_set_columns(&agg.group_expr)?,
+        &expand_grouping_sets(&agg.group_expr)?,
+    )?;
+
+    // The aggregate's output as Substrait orders it, which is what the
+    // expression below is written against.
+    let index_field = Field::new(GROUPING_SET_INDEX, DataType::Int32, false);
+    let substrait_output = DFSchema::from_unqualified_fields(
+        (0..grouping_id_index)
+            .chain(grouping_id_index + 1..schema.fields().len())
+            .map(|index| Arc::clone(schema.field(index)))
+            .chain(std::iter::once(Arc::new(index_field)))

Review Comment:
   Fixed: builds the temporary schema from qualified field specs 
(new_with_metadata) instead of from_unqualified_fields, so joined columns keep 
their qualifiers. Added a joined-column regression.



-- 
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]

Reply via email to