kosiew commented on code in PR #25527: URL: https://github.com/apache/datafusion/pull/25527#discussion_r4205221833
########## 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: This can shift by 64 when `group_count == 64`, which panics even when the ordinal is zero. Please handle the 64-column zero-ordinal case without shifting, reject nonzero ordinals that need more than 64 bits, and add producer and consumer boundary regressions. ########## datafusion/substrait/tests/cases/roundtrip_logical_plan.rs: ########## @@ -458,6 +475,107 @@ async fn aggregate_grouping_sets() -> Result<()> { Ok(()) } +/// `GROUPING()` reads the aggregate's `__grouping_id`, so it only survives a +/// round trip if the grouping set index Substrait carries is mapped back to it. +#[tokio::test] +async fn aggregate_grouping_sets_keep_grouping_function() -> Result<()> { + let ctx = create_context().await?; + // `(a)` first, then `(c)`: the bitmask is 1 then 2 while the index is 0 then + // 1, so a test that listed `(a, c)` first would pass either way + let sql = "SELECT a, c, GROUPING(a) AS ga, GROUPING(c) AS gc, avg(b) \ + FROM data GROUP BY GROUPING SETS ((a), (c), (a, c)) ORDER BY a, c, ga, gc"; + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let expected = DataFrame::new(ctx.state(), plan.clone()).collect().await?; + + let proto = to_substrait_plan(&plan, &ctx.state())?; + let plan2 = from_substrait_plan(&ctx.state(), &proto).await?; + let actual = DataFrame::new(ctx.state(), plan2).collect().await?; + + assert_eq!( + format!("{}", pretty_format_batches(&expected)?), + format!("{}", pretty_format_batches(&actual)?) + ); + // The values differ per set, so this is not vacuous + assert_snapshot!( + pretty_format_batches(&expected)?, + @r" + +---+------------+----+----+-------------+ + | a | c | ga | gc | avg(data.b) | + +---+------------+----+----+-------------+ + | 1 | 2020-01-01 | 0 | 0 | 2.000000 | + | 1 | | 0 | 1 | 2.000000 | + | 3 | 2020-01-01 | 0 | 0 | 4.500000 | + | 3 | | 0 | 1 | 4.500000 | + | | 2020-01-01 | 1 | 0 | 3.250000 | + +---+------------+----+----+-------------+ + " + ); + Ok(()) +} + +/// The same set listed twice is two sets in Substrait, and DataFusion separates +/// them with an ordinal packed above the `__grouping_id` bitmask. +#[tokio::test] +async fn aggregate_duplicate_grouping_sets() -> Result<()> { + let ctx = create_context().await?; + let sql = "SELECT a, GROUPING(a) AS ga, avg(b) \ + FROM data GROUP BY GROUPING SETS ((a), (a), ()) ORDER BY a, ga"; + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let expected = DataFrame::new(ctx.state(), plan.clone()).collect().await?; + + let proto = to_substrait_plan(&plan, &ctx.state())?; + let plan2 = from_substrait_plan(&ctx.state(), &proto).await?; + let actual = DataFrame::new(ctx.state(), plan2).collect().await?; + + assert_eq!( + format!("{}", pretty_format_batches(&expected)?), + format!("{}", pretty_format_batches(&actual)?) + ); + // Each occurrence of `(a)` keeps its own rows + assert_snapshot!( + pretty_format_batches(&expected)?, + @r" + +---+----+-------------+ + | a | ga | avg(data.b) | + +---+----+-------------+ + | 1 | 0 | 2.000000 | + | 1 | 0 | 2.000000 | + | 3 | 0 | 4.500000 | + | 3 | 0 | 4.500000 | + | | 1 | 3.250000 | + +---+----+-------------+ + " + ); + Ok(()) +} + +/// With more than eight grouping columns `__grouping_id` is a `UInt16` rather +/// than a `UInt8`, so the map back to it has to carry the wider literal. +#[tokio::test] +async fn aggregate_grouping_sets_wider_grouping_id() -> Result<()> { Review Comment: Suggestion: it would be useful to add a roundtrip with exactly eight distinct grouping expressions and a repeated set, so widening to `UInt16` is caused only by the ordinal bit. This would complement the existing duplicate-set and wider-ID tests. ########## 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: `from_unqualified_fields` drops qualifiers, so valid joined columns with the same base name can fail serialization as duplicate fields. Please preserve the qualified identities here, or use unique positional placeholders for this temporary schema, and add a joined-column regression. ########## datafusion/substrait/src/logical_plan/consumer/rel/aggregate_rel.rs: ########## @@ -160,11 +175,15 @@ fn reorder_grouping_set_output( measure_count ); } + let grouping_id = Expr::Column(grouping_id); + let grouping_id_type = grouping_id.get_type(schema)?; + let set_index = index_from_grouping_id(&grouping_id, &grouping_id_type, set_ids)? Review Comment: The fixed `grouping_set_index` alias can collide with a real user column of the same name on both the consumer and producer paths. Please generate a collision-free synthetic name and use that same chosen name for the producer reference, with regressions for both boundaries. -- 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]
