From 72f2f0956bfbe8ee6f07659bbe61f639166b9672 Mon Sep 17 00:00:00 2001 From: Niels Pardon Date: Tue, 10 Mar 2026 21:01:13 +0100 Subject: [PATCH 1/5] fix(core): use new aggregate grouping behavior in RelProtoConverter BREAKING CHANGE: changes RelProtoConverter to always output the new aggregate grouping proto structures Signed-off-by: Niels Pardon --- .../substrait/relation/RelProtoConverter.java | 21 +++++-- .../substrait/relation/AggregateRelTest.java | 59 +++++++++++++++++++ 2 files changed, 76 insertions(+), 4 deletions(-) diff --git a/core/src/main/java/io/substrait/relation/RelProtoConverter.java b/core/src/main/java/io/substrait/relation/RelProtoConverter.java index 02122b031..07766af0d 100644 --- a/core/src/main/java/io/substrait/relation/RelProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/RelProtoConverter.java @@ -155,12 +155,21 @@ private io.substrait.proto.Expression.FieldReference toProto(FieldReference fiel @Override public Rel visit(Aggregate aggregate, EmptyVisitationContext context) throws RuntimeException { - AggregateRel.Builder builder = + final List uniqueGroupingExpressions = + aggregate.getGroupings().stream() + .flatMap(g -> g.getExpressions().stream()) + .distinct() + .collect(Collectors.toList()); + + final AggregateRel.Builder builder = AggregateRel.newBuilder() .setInput(toProto(aggregate.getInput())) .setCommon(common(aggregate)) + .addAllGroupingExpressions(toProto(uniqueGroupingExpressions)) .addAllGroupings( - aggregate.getGroupings().stream().map(this::toProto).collect(Collectors.toList())) + aggregate.getGroupings().stream() + .map(g -> toProto(g, uniqueGroupingExpressions)) + .collect(Collectors.toList())) .addAllMeasures( aggregate.getMeasures().stream().map(this::toProto).collect(Collectors.toList())); @@ -203,9 +212,13 @@ private AggregateRel.Measure toProto(Aggregate.Measure measure) { return builder.build(); } - private AggregateRel.Grouping toProto(Aggregate.Grouping grouping) { + private AggregateRel.Grouping toProto( + Aggregate.Grouping grouping, List uniqueGroupingExpressions) { return AggregateRel.Grouping.newBuilder() - .addAllGroupingExpressions(toProto(grouping.getExpressions())) + .addAllExpressionReferences( + grouping.getExpressions().stream() + .map(e -> uniqueGroupingExpressions.indexOf(e)) + .collect(Collectors.toList())) .build(); } diff --git a/core/src/test/java/io/substrait/relation/AggregateRelTest.java b/core/src/test/java/io/substrait/relation/AggregateRelTest.java index 60857c472..50ee5079f 100644 --- a/core/src/test/java/io/substrait/relation/AggregateRelTest.java +++ b/core/src/test/java/io/substrait/relation/AggregateRelTest.java @@ -11,6 +11,7 @@ import io.substrait.proto.Plan; import io.substrait.proto.ReadRel; import io.substrait.proto.Rel; +import io.substrait.util.EmptyVisitationContext; import org.junit.jupiter.api.Test; class AggregateRelTest extends TestBase { @@ -90,6 +91,24 @@ void testDeprecatedGroupingExpressionConversion() { Aggregate agg = (Aggregate) resultRel; assertEquals(1, agg.getGroupings().size()); assertEquals(2, agg.getGroupings().get(0).getExpressions().size()); + + Rel roundtripRel = relProtoConverter.visit(agg, EmptyVisitationContext.INSTANCE); + assertTrue(roundtripRel.hasAggregate()); + + AggregateRel roundtripAggr = roundtripRel.getAggregate(); + assertEquals(1, roundtripAggr.getGroupingsCount(), "grouping count should be 1"); + assertEquals( + 2, + roundtripAggr.getGroupingExpressionsCount(), + "grouping expressions count of aggregate should be 2"); + assertEquals( + 0, + roundtripAggr.getGroupings(0).getGroupingExpressionsCount(), + "grouping expressions count of grouping should be 0"); + assertEquals( + 2, + roundtripAggr.getGroupings(0).getExpressionReferencesCount(), + "expression reference count of grouping should be 2"); } @Test @@ -127,6 +146,24 @@ void testAggregateWithSingleGrouping() { Aggregate agg = (Aggregate) resultRel; assertEquals(1, agg.getGroupings().size()); assertEquals(2, agg.getGroupings().get(0).getExpressions().size()); + + Rel roundtripRel = relProtoConverter.visit(agg, EmptyVisitationContext.INSTANCE); + assertTrue(roundtripRel.hasAggregate()); + + AggregateRel roundtripAgg = roundtripRel.getAggregate(); + assertEquals(1, roundtripAgg.getGroupingsCount(), "grouping count should be 1"); + assertEquals( + 2, + roundtripAgg.getGroupingExpressionsCount(), + "grouping expressions count of aggregate should be 2"); + assertEquals( + 0, + roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), + "grouping expressions count of grouping should be 0"); + assertEquals( + 2, + roundtripAgg.getGroupings(0).getExpressionReferencesCount(), + "expression reference count of grouping should be 2"); } @Test @@ -169,5 +206,27 @@ void testAggregateWithMultipleGroupings() { assertEquals(2, agg.getGroupings().size()); assertEquals(2, agg.getGroupings().get(0).getExpressions().size()); assertEquals(1, agg.getGroupings().get(1).getExpressions().size()); + + Rel roundtripRel = relProtoConverter.visit(agg, EmptyVisitationContext.INSTANCE); + assertTrue(roundtripRel.hasAggregate()); + + AggregateRel roundtripAgg = roundtripRel.getAggregate(); + assertEquals(2, roundtripAgg.getGroupingsCount(), "grouping count should be 2"); + assertEquals( + 2, + roundtripAgg.getGroupingExpressionsCount(), + "grouping expressions count of aggregate should be 2"); + assertEquals( + 0, + roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), + "grouping expressions count of grouping should be 0"); + assertEquals( + 2, + roundtripAgg.getGroupings(0).getExpressionReferencesCount(), + "expression reference count of grouping should be 2"); + assertEquals( + 1, + roundtripAgg.getGroupings(1).getExpressionReferencesCount(), + "expression reference count of grouping should be 2"); } } From 0ee022c0b7c176b8bb8ea207497ac0f4bc6ac218 Mon Sep 17 00:00:00 2001 From: Niels Pardon Date: Tue, 10 Mar 2026 21:19:34 +0100 Subject: [PATCH 2/5] fix: consistent variable naming Signed-off-by: Niels Pardon --- .../java/io/substrait/relation/AggregateRelTest.java | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/core/src/test/java/io/substrait/relation/AggregateRelTest.java b/core/src/test/java/io/substrait/relation/AggregateRelTest.java index 50ee5079f..c5e035394 100644 --- a/core/src/test/java/io/substrait/relation/AggregateRelTest.java +++ b/core/src/test/java/io/substrait/relation/AggregateRelTest.java @@ -95,19 +95,19 @@ void testDeprecatedGroupingExpressionConversion() { Rel roundtripRel = relProtoConverter.visit(agg, EmptyVisitationContext.INSTANCE); assertTrue(roundtripRel.hasAggregate()); - AggregateRel roundtripAggr = roundtripRel.getAggregate(); - assertEquals(1, roundtripAggr.getGroupingsCount(), "grouping count should be 1"); + AggregateRel roundtripAgg = roundtripRel.getAggregate(); + assertEquals(1, roundtripAgg.getGroupingsCount(), "grouping count should be 1"); assertEquals( 2, - roundtripAggr.getGroupingExpressionsCount(), + roundtripAgg.getGroupingExpressionsCount(), "grouping expressions count of aggregate should be 2"); assertEquals( 0, - roundtripAggr.getGroupings(0).getGroupingExpressionsCount(), + roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), "grouping expressions count of grouping should be 0"); assertEquals( 2, - roundtripAggr.getGroupings(0).getExpressionReferencesCount(), + roundtripAgg.getGroupings(0).getExpressionReferencesCount(), "expression reference count of grouping should be 2"); } From f6e7f5bde538238730336e008520b8023fe3bf06 Mon Sep 17 00:00:00 2001 From: Niels Pardon Date: Wed, 11 Mar 2026 20:01:38 +0100 Subject: [PATCH 3/5] fix: also output old structures Signed-off-by: Niels Pardon --- .../io/substrait/relation/RelProtoConverter.java | 2 ++ .../java/io/substrait/relation/AggregateRelTest.java | 12 ++++++------ 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/core/src/main/java/io/substrait/relation/RelProtoConverter.java b/core/src/main/java/io/substrait/relation/RelProtoConverter.java index 07766af0d..429d26661 100644 --- a/core/src/main/java/io/substrait/relation/RelProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/RelProtoConverter.java @@ -215,6 +215,8 @@ private AggregateRel.Measure toProto(Aggregate.Measure measure) { private AggregateRel.Grouping toProto( Aggregate.Grouping grouping, List uniqueGroupingExpressions) { return AggregateRel.Grouping.newBuilder() + .addAllGroupingExpressions( + grouping.getExpressions().stream().map(this::toProto).collect(Collectors.toList())) .addAllExpressionReferences( grouping.getExpressions().stream() .map(e -> uniqueGroupingExpressions.indexOf(e)) diff --git a/core/src/test/java/io/substrait/relation/AggregateRelTest.java b/core/src/test/java/io/substrait/relation/AggregateRelTest.java index c5e035394..390ecec7a 100644 --- a/core/src/test/java/io/substrait/relation/AggregateRelTest.java +++ b/core/src/test/java/io/substrait/relation/AggregateRelTest.java @@ -102,9 +102,9 @@ void testDeprecatedGroupingExpressionConversion() { roundtripAgg.getGroupingExpressionsCount(), "grouping expressions count of aggregate should be 2"); assertEquals( - 0, + 2, roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 0"); + "grouping expressions count of grouping should be 2"); assertEquals( 2, roundtripAgg.getGroupings(0).getExpressionReferencesCount(), @@ -157,9 +157,9 @@ void testAggregateWithSingleGrouping() { roundtripAgg.getGroupingExpressionsCount(), "grouping expressions count of aggregate should be 2"); assertEquals( - 0, + 2, roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 0"); + "grouping expressions count of grouping should be 2"); assertEquals( 2, roundtripAgg.getGroupings(0).getExpressionReferencesCount(), @@ -217,9 +217,9 @@ void testAggregateWithMultipleGroupings() { roundtripAgg.getGroupingExpressionsCount(), "grouping expressions count of aggregate should be 2"); assertEquals( - 0, + 2, roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 0"); + "grouping expressions count of grouping should be 2"); assertEquals( 2, roundtripAgg.getGroupings(0).getExpressionReferencesCount(), From 3408fce92188f71f7b7a680ecd69a30492d7086e Mon Sep 17 00:00:00 2001 From: Niels Pardon Date: Thu, 12 Mar 2026 21:40:43 +0100 Subject: [PATCH 4/5] fix: apply feedback Signed-off-by: Niels Pardon --- .../substrait/relation/AggregateRelTest.java | 99 +++++++++++-------- 1 file changed, 57 insertions(+), 42 deletions(-) diff --git a/core/src/test/java/io/substrait/relation/AggregateRelTest.java b/core/src/test/java/io/substrait/relation/AggregateRelTest.java index 390ecec7a..94202e4ce 100644 --- a/core/src/test/java/io/substrait/relation/AggregateRelTest.java +++ b/core/src/test/java/io/substrait/relation/AggregateRelTest.java @@ -12,6 +12,8 @@ import io.substrait.proto.ReadRel; import io.substrait.proto.Rel; import io.substrait.util.EmptyVisitationContext; +import java.util.List; +import java.util.stream.Collectors; import org.junit.jupiter.api.Test; class AggregateRelTest extends TestBase { @@ -58,6 +60,43 @@ public static io.substrait.proto.Expression createFieldReference(int col) { return Expression.newBuilder().setSelection(fieldRef1).build(); } + /** + * Helper method to extract expression references from an AggregateRel. + * + * @param aggregateRel the AggregateRel to extract expression references from + * @return a list of lists, where each inner list contains the expression reference indices for a + * grouping + */ + private static List> getExpressionReferences(AggregateRel aggregateRel) { + return aggregateRel.getGroupingsList().stream() + .map( + grouping -> + grouping.getExpressionReferencesList().stream().collect(Collectors.toList())) + .collect(Collectors.toList()); + } + + /** + * Helper method to extract deprecated grouping expressions from an AggregateRel. + * + * @param aggregateRel the AggregateRel to extract grouping expressions from + * @return a list of lists, where each inner list contains the grouping expressions for a grouping + */ + private static List> getGroupingExpressions(AggregateRel aggregateRel) { + return aggregateRel.getGroupingsList().stream() + .map(grouping -> grouping.getGroupingExpressionsList()) + .collect(Collectors.toList()); + } + + /** + * Helper method to extract aggregate-level grouping expressions from an AggregateRel. + * + * @param aggregateRel the AggregateRel to extract grouping expressions from + * @return a list of expressions at the aggregate level + */ + private static List getAggregateGroupingExpressions(AggregateRel aggregateRel) { + return aggregateRel.getGroupingExpressionsList(); + } + @Test void testDeprecatedGroupingExpressionConversion() { Expression col1Ref = createFieldReference(0); @@ -96,19 +135,12 @@ void testDeprecatedGroupingExpressionConversion() { assertTrue(roundtripRel.hasAggregate()); AggregateRel roundtripAgg = roundtripRel.getAggregate(); - assertEquals(1, roundtripAgg.getGroupingsCount(), "grouping count should be 1"); - assertEquals( - 2, - roundtripAgg.getGroupingExpressionsCount(), - "grouping expressions count of aggregate should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getExpressionReferencesCount(), - "expression reference count of grouping should be 2"); + // Verify new expression_references structure + assertEquals(List.of(List.of(0, 1)), getExpressionReferences(roundtripAgg)); + // Verify backward compatibility: deprecated grouping_expressions field is also populated + assertEquals(List.of(List.of(col1Ref, col2Ref)), getGroupingExpressions(roundtripAgg)); + // Verify aggregate-level grouping_expressions field is populated + assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); } @Test @@ -151,19 +183,12 @@ void testAggregateWithSingleGrouping() { assertTrue(roundtripRel.hasAggregate()); AggregateRel roundtripAgg = roundtripRel.getAggregate(); - assertEquals(1, roundtripAgg.getGroupingsCount(), "grouping count should be 1"); - assertEquals( - 2, - roundtripAgg.getGroupingExpressionsCount(), - "grouping expressions count of aggregate should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getExpressionReferencesCount(), - "expression reference count of grouping should be 2"); + // Verify new expression_references structure + assertEquals(List.of(List.of(0, 1)), getExpressionReferences(roundtripAgg)); + // Verify backward compatibility: deprecated grouping_expressions field is also populated + assertEquals(List.of(List.of(col1Ref, col2Ref)), getGroupingExpressions(roundtripAgg)); + // Verify aggregate-level grouping_expressions field is populated + assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); } @Test @@ -211,22 +236,12 @@ void testAggregateWithMultipleGroupings() { assertTrue(roundtripRel.hasAggregate()); AggregateRel roundtripAgg = roundtripRel.getAggregate(); - assertEquals(2, roundtripAgg.getGroupingsCount(), "grouping count should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupingExpressionsCount(), - "grouping expressions count of aggregate should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getGroupingExpressionsCount(), - "grouping expressions count of grouping should be 2"); - assertEquals( - 2, - roundtripAgg.getGroupings(0).getExpressionReferencesCount(), - "expression reference count of grouping should be 2"); + // Verify new expression_references structure + assertEquals(List.of(List.of(0, 1), List.of(1)), getExpressionReferences(roundtripAgg)); + // Verify backward compatibility: deprecated grouping_expressions field is also populated assertEquals( - 1, - roundtripAgg.getGroupings(1).getExpressionReferencesCount(), - "expression reference count of grouping should be 2"); + List.of(List.of(col1Ref, col2Ref), List.of(col2Ref)), getGroupingExpressions(roundtripAgg)); + // Verify aggregate-level grouping_expressions field is populated + assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); } } From ffba1384da5065fb14f7df689c04cc77bd0a8e0e Mon Sep 17 00:00:00 2001 From: Niels Pardon Date: Mon, 16 Mar 2026 10:06:23 +0100 Subject: [PATCH 5/5] fix: apply feedback Signed-off-by: Niels Pardon --- .../substrait/relation/AggregateRelTest.java | 48 ++++++++++++------- 1 file changed, 32 insertions(+), 16 deletions(-) diff --git a/core/src/test/java/io/substrait/relation/AggregateRelTest.java b/core/src/test/java/io/substrait/relation/AggregateRelTest.java index 94202e4ce..07f1e25cd 100644 --- a/core/src/test/java/io/substrait/relation/AggregateRelTest.java +++ b/core/src/test/java/io/substrait/relation/AggregateRelTest.java @@ -69,9 +69,7 @@ public static io.substrait.proto.Expression createFieldReference(int col) { */ private static List> getExpressionReferences(AggregateRel aggregateRel) { return aggregateRel.getGroupingsList().stream() - .map( - grouping -> - grouping.getExpressionReferencesList().stream().collect(Collectors.toList())) + .map(AggregateRel.Grouping::getExpressionReferencesList) .collect(Collectors.toList()); } @@ -87,16 +85,6 @@ private static List> getGroupingExpressions(AggregateRel aggreg .collect(Collectors.toList()); } - /** - * Helper method to extract aggregate-level grouping expressions from an AggregateRel. - * - * @param aggregateRel the AggregateRel to extract grouping expressions from - * @return a list of expressions at the aggregate level - */ - private static List getAggregateGroupingExpressions(AggregateRel aggregateRel) { - return aggregateRel.getGroupingExpressionsList(); - } - @Test void testDeprecatedGroupingExpressionConversion() { Expression col1Ref = createFieldReference(0); @@ -140,7 +128,7 @@ void testDeprecatedGroupingExpressionConversion() { // Verify backward compatibility: deprecated grouping_expressions field is also populated assertEquals(List.of(List.of(col1Ref, col2Ref)), getGroupingExpressions(roundtripAgg)); // Verify aggregate-level grouping_expressions field is populated - assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); + assertEquals(List.of(col1Ref, col2Ref), roundtripAgg.getGroupingExpressionsList()); } @Test @@ -188,7 +176,7 @@ void testAggregateWithSingleGrouping() { // Verify backward compatibility: deprecated grouping_expressions field is also populated assertEquals(List.of(List.of(col1Ref, col2Ref)), getGroupingExpressions(roundtripAgg)); // Verify aggregate-level grouping_expressions field is populated - assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); + assertEquals(List.of(col1Ref, col2Ref), roundtripAgg.getGroupingExpressionsList()); } @Test @@ -242,6 +230,34 @@ void testAggregateWithMultipleGroupings() { assertEquals( List.of(List.of(col1Ref, col2Ref), List.of(col2Ref)), getGroupingExpressions(roundtripAgg)); // Verify aggregate-level grouping_expressions field is populated - assertEquals(List.of(col1Ref, col2Ref), getAggregateGroupingExpressions(roundtripAgg)); + assertEquals(List.of(col1Ref, col2Ref), roundtripAgg.getGroupingExpressionsList()); + } + + /** + * Tests deduplication of non-trivial grouping expressions by equals() (not identity), and that an + * empty grouping set is handled correctly alongside non-empty ones. + */ + @Test + void testGroupingExpressionDeduplicationAndEmptyGroupingSet() { + NamedScan input = sb.namedScan(List.of("t"), List.of("a", "b"), List.of(R.I32, R.I32)); + + // Two independently constructed add(a, b) expressions — equal but not same instance. + io.substrait.expression.Expression addAB1 = + sb.add(sb.fieldReference(input, 0), sb.fieldReference(input, 1)); + io.substrait.expression.Expression addAB2 = + sb.add(sb.fieldReference(input, 0), sb.fieldReference(input, 1)); + + Aggregate agg = + Aggregate.builder() + .input(input) + .addGroupings(sb.grouping(addAB1), sb.grouping(addAB2), sb.grouping()) + .build(); + + Rel roundtripRel = relProtoConverter.visit(agg, EmptyVisitationContext.INSTANCE); + AggregateRel result = roundtripRel.getAggregate(); + + assertEquals( + List.of(expressionProtoConverter.toProto(addAB1)), result.getGroupingExpressionsList()); + assertEquals(List.of(List.of(0), List.of(0), List.of()), getExpressionReferences(result)); } }