From c78c66a2693ac8e837aa030695658056ab8ae664 Mon Sep 17 00:00:00 2001 From: Sandy Gao Date: Tue, 15 Jun 2021 22:25:06 -0400 Subject: [PATCH] Capture groupBy field in InMemoryAggregationBuilder --- .../aggregation/builder/InMemoryAggregationBuilder.java | 5 ++++- .../java/io/prestosql/spi/snapshot/SnapshotTestUtil.java | 5 ++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/presto-main/src/main/java/io/prestosql/operator/aggregation/builder/InMemoryAggregationBuilder.java b/presto-main/src/main/java/io/prestosql/operator/aggregation/builder/InMemoryAggregationBuilder.java index f9f182c89..81f4cff9b 100644 --- a/presto-main/src/main/java/io/prestosql/operator/aggregation/builder/InMemoryAggregationBuilder.java +++ b/presto-main/src/main/java/io/prestosql/operator/aggregation/builder/InMemoryAggregationBuilder.java @@ -51,7 +51,7 @@ import static io.prestosql.operator.GroupByHash.createGroupByHash; import static io.prestosql.operator.GroupBySort.createGroupBySort; import static java.util.Objects.requireNonNull; -@RestorableConfig(uncapturedFields = {"groupBy", "updateMemory"}) +@RestorableConfig(uncapturedFields = {"updateMemory"}) public abstract class InMemoryAggregationBuilder implements AggregationBuilder, Restorable { @@ -364,6 +364,7 @@ public abstract class InMemoryAggregationBuilder public Object capture(BlockEncodingSerdeProvider serdeProvider) { InMemoryAggregationBuilderState myState = new InMemoryAggregationBuilderState(); + myState.groupBy = groupBy.capture(serdeProvider); List aggregators = new ArrayList<>(); for (Aggregator aggregator : this.aggregators) { aggregators.add(aggregator.capture(serdeProvider)); @@ -377,6 +378,7 @@ public abstract class InMemoryAggregationBuilder public void restore(Object state, BlockEncodingSerdeProvider serdeProvider) { InMemoryAggregationBuilderState myState = (InMemoryAggregationBuilderState) state; + this.groupBy.restore(myState.groupBy, serdeProvider); for (int i = 0; i < this.aggregators.size(); i++) { this.aggregators.get(i).restore(myState.aggregators.get(i), serdeProvider); } @@ -392,6 +394,7 @@ public abstract class InMemoryAggregationBuilder private static class InMemoryAggregationBuilderState implements Serializable { + private Object groupBy; private List aggregators; private boolean full; } diff --git a/presto-spi/src/test/java/io/prestosql/spi/snapshot/SnapshotTestUtil.java b/presto-spi/src/test/java/io/prestosql/spi/snapshot/SnapshotTestUtil.java index 9dc1a459e..0351bc1e5 100644 --- a/presto-spi/src/test/java/io/prestosql/spi/snapshot/SnapshotTestUtil.java +++ b/presto-spi/src/test/java/io/prestosql/spi/snapshot/SnapshotTestUtil.java @@ -107,7 +107,10 @@ public final class SnapshotTestUtil for (Field field : obj.getClass().getDeclaredFields()) { field.setAccessible(true); try { - result.put(field.getName(), toFullSnapshotMapping(field.get(obj))); + // Too complex to compare groupBy content + if (!field.getName().equals("groupBy")) { + result.put(field.getName(), toFullSnapshotMapping(field.get(obj))); + } } catch (IllegalAccessException e) { e.printStackTrace();