From 0576c499617c43f5014d599dad3eb8b3f100aaaf Mon Sep 17 00:00:00 2001 From: Kevin Wan Date: Thu, 18 Nov 2021 14:48:06 -0500 Subject: [PATCH] Track memory usage in SingleInputSnapshotState --- .../prestosql/benchmark/HandTpchQuery1.java | 8 +++++ .../operator/AggregationOperator.java | 3 ++ .../operator/AssignUniqueIdOperator.java | 8 +++++ .../operator/DistinctLimitOperator.java | 8 +++++ .../operator/DynamicFilterSourceOperator.java | 8 +++++ .../operator/EnforceSingleRowOperator.java | 8 +++++ .../operator/ExplainAnalyzeOperator.java | 8 +++++ .../operator/FilterAndProjectOperator.java | 8 +++++ .../prestosql/operator/GroupIdOperator.java | 8 +++++ .../operator/HashAggregationOperator.java | 3 ++ .../operator/HashBuilderOperator.java | 5 +++ .../operator/HashSemiJoinOperator.java | 8 +++++ .../io/prestosql/operator/LimitOperator.java | 8 +++++ .../operator/LookupJoinOperator.java | 5 +++ .../operator/LookupOuterOperator.java | 3 ++ .../operator/MarkDistinctOperator.java | 8 +++++ .../operator/NestedLoopBuildOperator.java | 8 +++++ .../operator/NestedLoopJoinOperator.java | 3 ++ .../java/io/prestosql/operator/Operator.java | 9 ++++++ .../prestosql/operator/OperatorContext.java | 11 +++++++ .../prestosql/operator/OrderByOperator.java | 3 ++ .../operator/PartitionedOutputOperator.java | 8 +++++ .../prestosql/operator/RowNumberOperator.java | 8 +++++ .../operator/SetBuilderOperator.java | 8 +++++ .../operator/SortAggregationOperator.java | 3 ++ .../operator/SpatialIndexBuilderOperator.java | 3 ++ .../operator/SpatialJoinOperator.java | 4 ++- .../operator/StatisticsWriterOperator.java | 8 +++++ .../StreamingAggregationOperator.java | 8 +++++ .../operator/TableFinishOperator.java | 3 ++ .../operator/TableWriterOperator.java | 3 ++ .../operator/TaskOutputOperator.java | 8 +++++ .../operator/TopNRankingNumberOperator.java | 8 +++++ .../io/prestosql/operator/ValuesOperator.java | 8 +++++ .../io/prestosql/operator/WindowOperator.java | 3 ++ .../WorkProcessorOperatorAdapter.java | 3 ++ .../CrossRegionDynamicFilterOperator.java | 8 +++++ .../exchange/LocalExchangeSinkOperator.java | 3 ++ .../operator/unnest/UnnestOperator.java | 8 +++++ .../snapshot/SingleInputSnapshotState.java | 31 ++++++++++++++++--- .../execution/TestSqlStageExecution.java | 1 + .../operator/TestOrderByOperator.java | 6 ++-- .../operator/TestWindowOperator.java | 6 ++-- .../snapshot/TestMultiInputSnapshotState.java | 6 ++++ .../TestSingleInputSnapshotState.java | 24 +++++++++++--- .../io/prestosql/spi/snapshot/Restorable.java | 10 ++++++ 46 files changed, 315 insertions(+), 15 deletions(-) diff --git a/presto-benchmark/src/main/java/io/prestosql/benchmark/HandTpchQuery1.java b/presto-benchmark/src/main/java/io/prestosql/benchmark/HandTpchQuery1.java index eb9ca9f13..1882e077d 100644 --- a/presto-benchmark/src/main/java/io/prestosql/benchmark/HandTpchQuery1.java +++ b/presto-benchmark/src/main/java/io/prestosql/benchmark/HandTpchQuery1.java @@ -204,6 +204,14 @@ public class HandTpchQuery1 finishing = true; } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public boolean isFinished() { diff --git a/presto-main/src/main/java/io/prestosql/operator/AggregationOperator.java b/presto-main/src/main/java/io/prestosql/operator/AggregationOperator.java index 9b002cfee..4e8aa28f7 100644 --- a/presto-main/src/main/java/io/prestosql/operator/AggregationOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/AggregationOperator.java @@ -135,6 +135,9 @@ public class AggregationOperator { userMemoryContext.setBytes(0); systemMemoryContext.close(); + if (snapshotState != null) { + snapshotState.close(); + } } @Override diff --git a/presto-main/src/main/java/io/prestosql/operator/AssignUniqueIdOperator.java b/presto-main/src/main/java/io/prestosql/operator/AssignUniqueIdOperator.java index 81effb918..695bd4457 100644 --- a/presto-main/src/main/java/io/prestosql/operator/AssignUniqueIdOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/AssignUniqueIdOperator.java @@ -196,6 +196,14 @@ public class AssignUniqueIdOperator return block.build(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/DistinctLimitOperator.java b/presto-main/src/main/java/io/prestosql/operator/DistinctLimitOperator.java index 90564f935..512c724e7 100644 --- a/presto-main/src/main/java/io/prestosql/operator/DistinctLimitOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/DistinctLimitOperator.java @@ -289,6 +289,14 @@ public class DistinctLimitOperator return groupByHash.getCapacity(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/DynamicFilterSourceOperator.java b/presto-main/src/main/java/io/prestosql/operator/DynamicFilterSourceOperator.java index 490e3aaca..c0c7e0109 100644 --- a/presto-main/src/main/java/io/prestosql/operator/DynamicFilterSourceOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/DynamicFilterSourceOperator.java @@ -231,6 +231,14 @@ public class DynamicFilterSourceOperator } } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/EnforceSingleRowOperator.java b/presto-main/src/main/java/io/prestosql/operator/EnforceSingleRowOperator.java index 2a8cd2200..d94b573ca 100644 --- a/presto-main/src/main/java/io/prestosql/operator/EnforceSingleRowOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/EnforceSingleRowOperator.java @@ -161,6 +161,14 @@ public class EnforceSingleRowOperator return snapshotState.nextMarker(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/ExplainAnalyzeOperator.java b/presto-main/src/main/java/io/prestosql/operator/ExplainAnalyzeOperator.java index b306058a3..91a9a60fa 100644 --- a/presto-main/src/main/java/io/prestosql/operator/ExplainAnalyzeOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/ExplainAnalyzeOperator.java @@ -226,6 +226,14 @@ public class ExplainAnalyzeOperator } } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/FilterAndProjectOperator.java b/presto-main/src/main/java/io/prestosql/operator/FilterAndProjectOperator.java index 475c8eb00..562b729f1 100644 --- a/presto-main/src/main/java/io/prestosql/operator/FilterAndProjectOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/FilterAndProjectOperator.java @@ -134,6 +134,14 @@ public class FilterAndProjectOperator return snapshotState.nextMarker(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/GroupIdOperator.java b/presto-main/src/main/java/io/prestosql/operator/GroupIdOperator.java index f9105156a..0325db166 100644 --- a/presto-main/src/main/java/io/prestosql/operator/GroupIdOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/GroupIdOperator.java @@ -233,6 +233,14 @@ public class GroupIdOperator return outputPage; } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/HashAggregationOperator.java b/presto-main/src/main/java/io/prestosql/operator/HashAggregationOperator.java index 92eba7d5e..72bb72015 100644 --- a/presto-main/src/main/java/io/prestosql/operator/HashAggregationOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/HashAggregationOperator.java @@ -397,6 +397,9 @@ public class HashAggregationOperator public void close() { closeAggregationBuilder(); + if (snapshotState != null) { + snapshotState.close(); + } } protected void closeAggregationBuilder() diff --git a/presto-main/src/main/java/io/prestosql/operator/HashBuilderOperator.java b/presto-main/src/main/java/io/prestosql/operator/HashBuilderOperator.java index 1bb652b34..41fc83a26 100644 --- a/presto-main/src/main/java/io/prestosql/operator/HashBuilderOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/HashBuilderOperator.java @@ -697,6 +697,11 @@ public class HashBuilderOperator spiller.ifPresent(closer::register); closer.register(() -> localUserMemoryContext.setBytes(0)); closer.register(() -> localRevocableMemoryContext.setBytes(0)); + closer.register(() -> { + if (snapshotState != null) { + snapshotState.close(); + } + }); } catch (IOException e) { throw new RuntimeException(e); diff --git a/presto-main/src/main/java/io/prestosql/operator/HashSemiJoinOperator.java b/presto-main/src/main/java/io/prestosql/operator/HashSemiJoinOperator.java index 90eebb271..97958589a 100644 --- a/presto-main/src/main/java/io/prestosql/operator/HashSemiJoinOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/HashSemiJoinOperator.java @@ -237,6 +237,14 @@ public class HashSemiJoinOperator return snapshotState.nextMarker(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/LimitOperator.java b/presto-main/src/main/java/io/prestosql/operator/LimitOperator.java index 404ae87ba..42b024ebe 100644 --- a/presto-main/src/main/java/io/prestosql/operator/LimitOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/LimitOperator.java @@ -157,6 +157,14 @@ public class LimitOperator return snapshotState.nextMarker(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/LookupJoinOperator.java b/presto-main/src/main/java/io/prestosql/operator/LookupJoinOperator.java index 75afbf8dd..bdd13015e 100644 --- a/presto-main/src/main/java/io/prestosql/operator/LookupJoinOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/LookupJoinOperator.java @@ -589,6 +589,11 @@ public class LookupJoinOperator closer.register(pageBuilder::reset); closer.register(() -> Optional.ofNullable(lookupSourceProvider).ifPresent(LookupSourceProvider::close)); + closer.register(() -> { + if (snapshotState != null) { + snapshotState.close(); + } + }); spiller.ifPresent(closer::register); } catch (IOException e) { diff --git a/presto-main/src/main/java/io/prestosql/operator/LookupOuterOperator.java b/presto-main/src/main/java/io/prestosql/operator/LookupOuterOperator.java index 92c0c7a11..25573cd23 100644 --- a/presto-main/src/main/java/io/prestosql/operator/LookupOuterOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/LookupOuterOperator.java @@ -290,6 +290,9 @@ public class LookupOuterOperator if (closed) { return; } + if (snapshotState != null) { + snapshotState.close(); + } closed = true; pageBuilder.reset(); onClose.run(); diff --git a/presto-main/src/main/java/io/prestosql/operator/MarkDistinctOperator.java b/presto-main/src/main/java/io/prestosql/operator/MarkDistinctOperator.java index eb77f04f7..ed246b78b 100644 --- a/presto-main/src/main/java/io/prestosql/operator/MarkDistinctOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/MarkDistinctOperator.java @@ -229,6 +229,14 @@ public class MarkDistinctOperator return markDistinctHash.getCapacity(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/NestedLoopBuildOperator.java b/presto-main/src/main/java/io/prestosql/operator/NestedLoopBuildOperator.java index 57d37d10c..a154cad57 100644 --- a/presto-main/src/main/java/io/prestosql/operator/NestedLoopBuildOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/NestedLoopBuildOperator.java @@ -153,6 +153,14 @@ public class NestedLoopBuildOperator operatorContext.recordOutput(page.getSizeInBytes(), page.getPositionCount()); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/NestedLoopJoinOperator.java b/presto-main/src/main/java/io/prestosql/operator/NestedLoopJoinOperator.java index 9c06106d9..4bfaa4e67 100644 --- a/presto-main/src/main/java/io/prestosql/operator/NestedLoopJoinOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/NestedLoopJoinOperator.java @@ -259,6 +259,9 @@ public class NestedLoopJoinOperator if (closed) { return; } + if (snapshotState != null) { + snapshotState.close(); + } closed = true; // `afterClose` must be run last. afterClose.run(); diff --git a/presto-main/src/main/java/io/prestosql/operator/Operator.java b/presto-main/src/main/java/io/prestosql/operator/Operator.java index 33fc7f22a..7f3da7013 100644 --- a/presto-main/src/main/java/io/prestosql/operator/Operator.java +++ b/presto-main/src/main/java/io/prestosql/operator/Operator.java @@ -128,4 +128,13 @@ public interface Operator { close(); } + + /** + * Estimate the total memory used in an operator state by copying the total memory used by this operator + */ + @Override + default long getUsedMemory() + { + return getOperatorContext().getTotalMemoryBytes(); + } } diff --git a/presto-main/src/main/java/io/prestosql/operator/OperatorContext.java b/presto-main/src/main/java/io/prestosql/operator/OperatorContext.java index 564c46421..2b139f5af 100644 --- a/presto-main/src/main/java/io/prestosql/operator/OperatorContext.java +++ b/presto-main/src/main/java/io/prestosql/operator/OperatorContext.java @@ -267,6 +267,12 @@ public class OperatorContext return new InternalLocalMemoryContext(operatorMemoryContext.newSystemMemoryContext(allocationTag), memoryFuture, this::updatePeakMemoryReservations, true); } + // caller should close this context as it's a new context + public LocalMemoryContext newLocalUserMemoryContext(String allocationTag) + { + return new InternalLocalMemoryContext(operatorMemoryContext.newUserMemoryContext(allocationTag), memoryFuture, this::updatePeakMemoryReservations, true); + } + // caller shouldn't close this context as it's managed by the OperatorContext public LocalMemoryContext localUserMemoryContext() { @@ -339,6 +345,11 @@ public class OperatorContext return operatorMemoryContext.getRevocableMemory(); } + public long getTotalMemoryBytes() + { + return operatorMemoryContext.getUserMemory() + operatorMemoryContext.getSystemMemory(); + } + private static void updateMemoryFuture(ListenableFuture memoryPoolFuture, AtomicReference> targetFutureReference) { if (!memoryPoolFuture.isDone()) { diff --git a/presto-main/src/main/java/io/prestosql/operator/OrderByOperator.java b/presto-main/src/main/java/io/prestosql/operator/OrderByOperator.java index 5548c2fa2..4b305e1f3 100644 --- a/presto-main/src/main/java/io/prestosql/operator/OrderByOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/OrderByOperator.java @@ -418,6 +418,9 @@ public class OrderByOperator @Override public void close() { + if (snapshotState != null) { + snapshotState.close(); + } pageIndex.clear(); sortedPages = null; spiller.ifPresent(Spiller::close); diff --git a/presto-main/src/main/java/io/prestosql/operator/PartitionedOutputOperator.java b/presto-main/src/main/java/io/prestosql/operator/PartitionedOutputOperator.java index f8cad97f7..aedaf6255 100644 --- a/presto-main/src/main/java/io/prestosql/operator/PartitionedOutputOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/PartitionedOutputOperator.java @@ -612,6 +612,14 @@ public class PartitionedOutputOperator } } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/RowNumberOperator.java b/presto-main/src/main/java/io/prestosql/operator/RowNumberOperator.java index 76920cca7..de302fd33 100644 --- a/presto-main/src/main/java/io/prestosql/operator/RowNumberOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/RowNumberOperator.java @@ -393,6 +393,14 @@ public class RowNumberOperator return groupByHash.map(GroupByHash::getCapacity).orElse(0); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/SetBuilderOperator.java b/presto-main/src/main/java/io/prestosql/operator/SetBuilderOperator.java index 2ea45d9c7..3401df743 100644 --- a/presto-main/src/main/java/io/prestosql/operator/SetBuilderOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/SetBuilderOperator.java @@ -236,6 +236,14 @@ public class SetBuilderOperator return channelSetBuilder.getCapacity(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/SortAggregationOperator.java b/presto-main/src/main/java/io/prestosql/operator/SortAggregationOperator.java index 00f1e1253..daeb68eb2 100644 --- a/presto-main/src/main/java/io/prestosql/operator/SortAggregationOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/SortAggregationOperator.java @@ -367,6 +367,9 @@ public class SortAggregationOperator public void close() { closeAggregationBuilder(); + if (snapshotState != null) { + snapshotState.close(); + } } protected void closeAggregationBuilder() diff --git a/presto-main/src/main/java/io/prestosql/operator/SpatialIndexBuilderOperator.java b/presto-main/src/main/java/io/prestosql/operator/SpatialIndexBuilderOperator.java index 7166f9cff..9d9228841 100644 --- a/presto-main/src/main/java/io/prestosql/operator/SpatialIndexBuilderOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/SpatialIndexBuilderOperator.java @@ -263,6 +263,9 @@ public class SpatialIndexBuilderOperator @Override public void close() { + if (snapshotState != null) { + snapshotState.close(); + } index.clear(); localUserMemoryContext.setBytes(index.getEstimatedSize().toBytes()); } diff --git a/presto-main/src/main/java/io/prestosql/operator/SpatialJoinOperator.java b/presto-main/src/main/java/io/prestosql/operator/SpatialJoinOperator.java index 3e813c253..23bc15bf3 100644 --- a/presto-main/src/main/java/io/prestosql/operator/SpatialJoinOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/SpatialJoinOperator.java @@ -368,7 +368,9 @@ public class SpatialJoinOperator return; } closed = true; - + if (snapshotState != null) { + snapshotState.close(); + } pagesSpatialIndexFuture = null; onClose.run(); } diff --git a/presto-main/src/main/java/io/prestosql/operator/StatisticsWriterOperator.java b/presto-main/src/main/java/io/prestosql/operator/StatisticsWriterOperator.java index bea3ee9cb..0272475cb 100644 --- a/presto-main/src/main/java/io/prestosql/operator/StatisticsWriterOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/StatisticsWriterOperator.java @@ -228,6 +228,14 @@ public class StatisticsWriterOperator void writeStatistics(Collection computedStatistics); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/StreamingAggregationOperator.java b/presto-main/src/main/java/io/prestosql/operator/StreamingAggregationOperator.java index 438cb5d30..68451b308 100644 --- a/presto-main/src/main/java/io/prestosql/operator/StreamingAggregationOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/StreamingAggregationOperator.java @@ -329,6 +329,14 @@ public class StreamingAggregationOperator return finishing && outputPages.isEmpty() && currentGroup == null && pageBuilder.isEmpty(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/TableFinishOperator.java b/presto-main/src/main/java/io/prestosql/operator/TableFinishOperator.java index f7d104322..b284e505a 100644 --- a/presto-main/src/main/java/io/prestosql/operator/TableFinishOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/TableFinishOperator.java @@ -369,6 +369,9 @@ public class TableFinishOperator public void close() throws Exception { + if (snapshotState != null) { + snapshotState.close(); + } statisticsAggregationOperator.close(); } diff --git a/presto-main/src/main/java/io/prestosql/operator/TableWriterOperator.java b/presto-main/src/main/java/io/prestosql/operator/TableWriterOperator.java index 0977e92ee..1d88daad6 100644 --- a/presto-main/src/main/java/io/prestosql/operator/TableWriterOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/TableWriterOperator.java @@ -383,6 +383,9 @@ public class TableWriterOperator public void close() throws Exception { + if (snapshotState != null) { + snapshotState.close(); + } closeImpl(pageSink::abort); } diff --git a/presto-main/src/main/java/io/prestosql/operator/TaskOutputOperator.java b/presto-main/src/main/java/io/prestosql/operator/TaskOutputOperator.java index 840575875..97dd7f5d0 100644 --- a/presto-main/src/main/java/io/prestosql/operator/TaskOutputOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/TaskOutputOperator.java @@ -211,6 +211,14 @@ public class TaskOutputOperator operatorContext.recordOutput(page.getSizeInBytes(), page.getPositionCount()); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/TopNRankingNumberOperator.java b/presto-main/src/main/java/io/prestosql/operator/TopNRankingNumberOperator.java index c4bf3cd44..4d49743b1 100644 --- a/presto-main/src/main/java/io/prestosql/operator/TopNRankingNumberOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/TopNRankingNumberOperator.java @@ -334,6 +334,14 @@ public class TopNRankingNumberOperator return types.build(); } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/ValuesOperator.java b/presto-main/src/main/java/io/prestosql/operator/ValuesOperator.java index 3e4ee2c3e..8a4601b7f 100644 --- a/presto-main/src/main/java/io/prestosql/operator/ValuesOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/ValuesOperator.java @@ -155,6 +155,14 @@ public class ValuesOperator return marker; } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/WindowOperator.java b/presto-main/src/main/java/io/prestosql/operator/WindowOperator.java index 7bcef28d5..3daabf9ff 100644 --- a/presto-main/src/main/java/io/prestosql/operator/WindowOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/WindowOperator.java @@ -1272,6 +1272,9 @@ public class WindowOperator @Override public void close() { + if (snapshotState != null) { + snapshotState.close(); + } driverWindowInfo.set(Optional.of(windowInfo.build())); spillablePagesToPagesIndexes.ifPresent(SpillablePagesToPagesIndexes::closeSpiller); } diff --git a/presto-main/src/main/java/io/prestosql/operator/WorkProcessorOperatorAdapter.java b/presto-main/src/main/java/io/prestosql/operator/WorkProcessorOperatorAdapter.java index 56b258914..ec96b2359 100644 --- a/presto-main/src/main/java/io/prestosql/operator/WorkProcessorOperatorAdapter.java +++ b/presto-main/src/main/java/io/prestosql/operator/WorkProcessorOperatorAdapter.java @@ -151,6 +151,9 @@ public class WorkProcessorOperatorAdapter public void close() throws Exception { + if (snapshotState != null) { + snapshotState.close(); + } workProcessorOperator.close(); } diff --git a/presto-main/src/main/java/io/prestosql/operator/dynamicfilter/CrossRegionDynamicFilterOperator.java b/presto-main/src/main/java/io/prestosql/operator/dynamicfilter/CrossRegionDynamicFilterOperator.java index 62aa83432..a8937f4b5 100644 --- a/presto-main/src/main/java/io/prestosql/operator/dynamicfilter/CrossRegionDynamicFilterOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/dynamicfilter/CrossRegionDynamicFilterOperator.java @@ -195,6 +195,14 @@ public class CrossRegionDynamicFilterOperator } } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/operator/exchange/LocalExchangeSinkOperator.java b/presto-main/src/main/java/io/prestosql/operator/exchange/LocalExchangeSinkOperator.java index e7b069d1b..512164a49 100644 --- a/presto-main/src/main/java/io/prestosql/operator/exchange/LocalExchangeSinkOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/exchange/LocalExchangeSinkOperator.java @@ -177,6 +177,9 @@ public class LocalExchangeSinkOperator @Override public void close() { + if (snapshotState != null) { + snapshotState.close(); + } finish(); } diff --git a/presto-main/src/main/java/io/prestosql/operator/unnest/UnnestOperator.java b/presto-main/src/main/java/io/prestosql/operator/unnest/UnnestOperator.java index d39360fdb..ff0cad2a1 100644 --- a/presto-main/src/main/java/io/prestosql/operator/unnest/UnnestOperator.java +++ b/presto-main/src/main/java/io/prestosql/operator/unnest/UnnestOperator.java @@ -336,6 +336,14 @@ public class UnnestOperator } } + @Override + public void close() + { + if (snapshotState != null) { + snapshotState.close(); + } + } + @Override public Object capture(BlockEncodingSerdeProvider serdeProvider) { diff --git a/presto-main/src/main/java/io/prestosql/snapshot/SingleInputSnapshotState.java b/presto-main/src/main/java/io/prestosql/snapshot/SingleInputSnapshotState.java index f6fdbfeb8..2684291e1 100644 --- a/presto-main/src/main/java/io/prestosql/snapshot/SingleInputSnapshotState.java +++ b/presto-main/src/main/java/io/prestosql/snapshot/SingleInputSnapshotState.java @@ -16,6 +16,7 @@ package io.prestosql.snapshot; import io.airlift.log.Logger; import io.hetu.core.transport.execution.buffer.PagesSerde; +import io.prestosql.memory.context.LocalMemoryContext; import io.prestosql.operator.Operator; import io.prestosql.operator.OperatorContext; import io.prestosql.spi.Page; @@ -35,6 +36,7 @@ import static java.util.Objects.requireNonNull; * This is a utility class used by non-source operators, which only receive inputs from a single source. * When an input is received from addInput(Page), the operator first calls the processPage() method to perform snapshot related processing. * When getOutput() is called, the operator calls the pollMarker() method to determine if a marker page needs to be returned. + * Any instance of this class created through forOperator needs to be closed */ public class SingleInputSnapshotState { @@ -52,6 +54,8 @@ public class SingleInputSnapshotState private final Function spillStateIdGenerator; // Markers to be returned to the restorable object. The "nextMarker" method polls this list. private final Queue markers = new LinkedList<>(); + // For recording snapshot memory usage + private final LocalMemoryContext snapshotMemoryContext; public static SingleInputSnapshotState forOperator(Operator operator, OperatorContext operatorContext) { @@ -60,14 +64,16 @@ public class SingleInputSnapshotState operatorContext.getDriverContext().getPipelineContext().getTaskContext().getSnapshotManager(), operatorContext.getDriverContext().getSerde(), snapshotId -> SnapshotStateId.forOperator(snapshotId, operatorContext), - snapshotId -> SnapshotStateId.forDriverComponent(snapshotId, operatorContext, operatorContext.getOperatorId() + "-spill")); + snapshotId -> SnapshotStateId.forDriverComponent(snapshotId, operatorContext, operatorContext.getOperatorId() + "-spill"), + operatorContext.newLocalUserMemoryContext(SingleInputSnapshotState.class.getSimpleName())); } SingleInputSnapshotState(Restorable restorable, - TaskSnapshotManager snapshotManager, - PagesSerde pagesSerde, - Function snapshotStateIdGenerator, - Function spillStateIdGenerator) + TaskSnapshotManager snapshotManager, + PagesSerde pagesSerde, + Function snapshotStateIdGenerator, + Function spillStateIdGenerator, + LocalMemoryContext snapshotMemoryContext) { this.restorable = requireNonNull(restorable, "restorable is null"); this.restorableId = String.format("%s (%s)", restorable.getClass().getSimpleName(), snapshotStateIdGenerator.apply(0L).getId()); @@ -75,6 +81,12 @@ public class SingleInputSnapshotState this.snapshotStateIdGenerator = requireNonNull(snapshotStateIdGenerator, "snapshotStateIdGenerator is null"); this.spillStateIdGenerator = requireNonNull(spillStateIdGenerator, "spillStateIdGenerator is null"); this.pagesSerde = pagesSerde; + this.snapshotMemoryContext = snapshotMemoryContext; + } + + public void close() + { + snapshotMemoryContext.close(); } /** @@ -154,6 +166,12 @@ public class SingleInputSnapshotState private void captureState(long snapshotId, boolean record) { SnapshotStateId componentId = snapshotStateIdGenerator.apply(snapshotId); + long stateMemory = restorable.getUsedMemory(); + if (!snapshotMemoryContext.trySetBytes(stateMemory)) { + LOG.warn("Insufficient memory on worker node to take snapshot"); + snapshotManager.failedToCapture(componentId); + return; + } try { if (restorable.supportsConsolidatedWrites()) { snapshotManager.storeConsolidatedState(componentId, restorable.capture(pagesSerde)); @@ -176,6 +194,9 @@ public class SingleInputSnapshotState LOG.warn(e, "Failed to capture and store snapshot state"); snapshotManager.failedToCapture(componentId); } + finally { + snapshotMemoryContext.setBytes(0); + } } public boolean hasMarker() diff --git a/presto-main/src/test/java/io/prestosql/execution/TestSqlStageExecution.java b/presto-main/src/test/java/io/prestosql/execution/TestSqlStageExecution.java index 5979121f1..3c75e6104 100644 --- a/presto-main/src/test/java/io/prestosql/execution/TestSqlStageExecution.java +++ b/presto-main/src/test/java/io/prestosql/execution/TestSqlStageExecution.java @@ -68,6 +68,7 @@ import static org.testng.Assert.assertFalse; import static org.testng.Assert.assertSame; import static org.testng.Assert.assertTrue; +@Test(singleThreaded = true) public class TestSqlStageExecution { private ExecutorService executor; diff --git a/presto-main/src/test/java/io/prestosql/operator/TestOrderByOperator.java b/presto-main/src/test/java/io/prestosql/operator/TestOrderByOperator.java index fca4eff97..609c01a8f 100644 --- a/presto-main/src/test/java/io/prestosql/operator/TestOrderByOperator.java +++ b/presto-main/src/test/java/io/prestosql/operator/TestOrderByOperator.java @@ -84,6 +84,8 @@ import static org.testng.Assert.assertTrue; @Test(singleThreaded = true) public class TestOrderByOperator { + private static final long defaultMemoryLimit = 1L << 28; + private ExecutorService executor; private ScheduledExecutorService scheduledExecutor; private DummySpillerFactory spillerFactory; @@ -404,7 +406,7 @@ public class TestOrderByOperator Optional.of(spillerFactory), new OrderingCompiler()); - DriverContext driverContext = createDriverContext(8, TEST_SNAPSHOT_SESSION); + DriverContext driverContext = createDriverContext(defaultMemoryLimit, TEST_SNAPSHOT_SESSION); driverContext.getPipelineContext().getTaskContext().getSnapshotManager().setTotalComponents(1); OrderByOperator orderByOperator = (OrderByOperator) operatorFactory.createOperator(driverContext); @@ -426,7 +428,7 @@ public class TestOrderByOperator } // Step5: assume the task is rescheduled due to failure and everything is re-constructed - driverContext = createDriverContext(8, TEST_SNAPSHOT_SESSION); + driverContext = createDriverContext(defaultMemoryLimit, TEST_SNAPSHOT_SESSION); operatorFactory = new OrderByOperatorFactory( 0, new PlanNodeId("test"), diff --git a/presto-main/src/test/java/io/prestosql/operator/TestWindowOperator.java b/presto-main/src/test/java/io/prestosql/operator/TestWindowOperator.java index bb78250d7..11be3be17 100644 --- a/presto-main/src/test/java/io/prestosql/operator/TestWindowOperator.java +++ b/presto-main/src/test/java/io/prestosql/operator/TestWindowOperator.java @@ -120,6 +120,8 @@ public class TestWindowOperator private static final List LEAD = ImmutableList.of( window(new ReflectionWindowFunctionSupplier<>("lead", VARCHAR, ImmutableList.of(VARCHAR, BIGINT, VARCHAR), LeadFunction.class), VARCHAR, UNBOUNDED_FRAME, 1, 3, 4)); + private static final long defaultMemoryLimit = 1L << 28; + private ExecutorService executor; private ScheduledExecutorService scheduledExecutor; private DummySpillerFactory spillerFactory; @@ -1223,7 +1225,7 @@ public class TestWindowOperator spillerFactory, new OrderingCompiler()); - DriverContext driverContext = createDriverContext(0, TEST_SNAPSHOT_SESSION); + DriverContext driverContext = createDriverContext(defaultMemoryLimit, TEST_SNAPSHOT_SESSION); WindowOperator windowOperator = (WindowOperator) operatorFactory.createOperator(driverContext); // Step1: add the first 2 pages @@ -1340,7 +1342,7 @@ public class TestWindowOperator ImmutableList.copyOf(new SortOrder[] {SortOrder.ASC_NULLS_LAST}), false); - DriverContext driverContext = createDriverContext(0, TEST_SNAPSHOT_SESSION); + DriverContext driverContext = createDriverContext(defaultMemoryLimit, TEST_SNAPSHOT_SESSION); WindowOperator windowOperator = (WindowOperator) operatorFactory.createOperator(driverContext); // Step1: add the first 2 pages diff --git a/presto-main/src/test/java/io/prestosql/snapshot/TestMultiInputSnapshotState.java b/presto-main/src/test/java/io/prestosql/snapshot/TestMultiInputSnapshotState.java index 8a735e992..2f201a4da 100644 --- a/presto-main/src/test/java/io/prestosql/snapshot/TestMultiInputSnapshotState.java +++ b/presto-main/src/test/java/io/prestosql/snapshot/TestMultiInputSnapshotState.java @@ -516,6 +516,12 @@ public class TestMultiInputSnapshotState return this.supportsConsolidatedWrites; } + @Override + public long getUsedMemory() + { + return 0; + } + public void setSupportsConsolidatedWrites(boolean supportsConsolidatedWrites) { this.supportsConsolidatedWrites = supportsConsolidatedWrites; diff --git a/presto-main/src/test/java/io/prestosql/snapshot/TestSingleInputSnapshotState.java b/presto-main/src/test/java/io/prestosql/snapshot/TestSingleInputSnapshotState.java index acbad2124..15c344538 100644 --- a/presto-main/src/test/java/io/prestosql/snapshot/TestSingleInputSnapshotState.java +++ b/presto-main/src/test/java/io/prestosql/snapshot/TestSingleInputSnapshotState.java @@ -16,6 +16,7 @@ package io.prestosql.snapshot; import com.google.common.collect.ImmutableList; import io.prestosql.execution.TaskId; +import io.prestosql.memory.context.LocalMemoryContext; import io.prestosql.operator.DriverContext; import io.prestosql.operator.Operator; import io.prestosql.operator.OperatorContext; @@ -39,6 +40,7 @@ import static io.prestosql.SessionTestUtils.TEST_SNAPSHOT_SESSION; import static io.prestosql.testing.TestingTaskContext.createTaskContext; import static java.util.concurrent.Executors.newScheduledThreadPool; import static org.mockito.Matchers.anyBoolean; +import static org.mockito.Matchers.anyLong; import static org.mockito.Matchers.anyObject; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; @@ -64,15 +66,18 @@ public class TestSingleInputSnapshotState private TaskSnapshotManager snapshotManager; private TestingRestorable restorable; private SingleInputSnapshotState state; + private LocalMemoryContext snapshotMemoryContext; @BeforeMethod public void setup() throws Exception { snapshotManager = mock(TaskSnapshotManager.class); + snapshotMemoryContext = mock(LocalMemoryContext.class); + when(snapshotMemoryContext.trySetBytes(anyLong())).thenReturn(true); restorable = new TestingRestorable(); restorable.state = 100; - state = new SingleInputSnapshotState(restorable, snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, TestSingleInputSnapshotState::createSnapshotStateId); + state = new SingleInputSnapshotState(restorable, snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, TestSingleInputSnapshotState::createSnapshotStateId, snapshotMemoryContext); } private boolean processPage(Page page) @@ -199,7 +204,7 @@ public class TestSingleInputSnapshotState public void testResumeBacktrack() throws Exception { - SingleInputSnapshotState state = new SingleInputSnapshotState(restorable, snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, TestSingleInputSnapshotState::createSnapshotStateId); + SingleInputSnapshotState state = new SingleInputSnapshotState(restorable, snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, TestSingleInputSnapshotState::createSnapshotStateId, snapshotMemoryContext); state.processPage(regularPage); restorable.state++; int saved1 = restorable.state; @@ -222,7 +227,8 @@ public class TestSingleInputSnapshotState snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, - TestSingleInputSnapshotState::createSnapshotStateId); + TestSingleInputSnapshotState::createSnapshotStateId, + snapshotMemoryContext); state.processPage(marker1); when(snapshotManager.loadState(anyObject())).thenReturn(Optional.of(1)); when(snapshotManager.loadFile(anyObject(), anyObject())) @@ -250,7 +256,8 @@ public class TestSingleInputSnapshotState snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, - TestSingleInputSnapshotState::createSnapshotStateId); + TestSingleInputSnapshotState::createSnapshotStateId, + snapshotMemoryContext); state.processPage(marker1); when(snapshotManager.loadConsolidatedState(anyObject())).thenReturn(Optional.of(0)); state.processPage(resume1); @@ -271,7 +278,8 @@ public class TestSingleInputSnapshotState snapshotManager, null, TestSingleInputSnapshotState::createSnapshotStateId, - TestSingleInputSnapshotState::createSnapshotStateId); + TestSingleInputSnapshotState::createSnapshotStateId, + snapshotMemoryContext); state.processPage(marker1); when(snapshotManager.loadState(anyObject())).thenReturn(Optional.of(0)); state.processPage(resume1); @@ -306,6 +314,12 @@ public class TestSingleInputSnapshotState return this.supportsConsolidatedWrites; } + @Override + public long getUsedMemory() + { + return 0; + } + public void setSupportsConsolidatedWrites(boolean supportsConsolidatedWrites) { this.supportsConsolidatedWrites = supportsConsolidatedWrites; diff --git a/presto-spi/src/main/java/io/prestosql/spi/snapshot/Restorable.java b/presto-spi/src/main/java/io/prestosql/spi/snapshot/Restorable.java index 82f7c254b..b38091dc7 100644 --- a/presto-spi/src/main/java/io/prestosql/spi/snapshot/Restorable.java +++ b/presto-spi/src/main/java/io/prestosql/spi/snapshot/Restorable.java @@ -50,4 +50,14 @@ public interface Restorable { return true; } + + /** + * Finds the memory used to capture this object in a snapshot + * + * @return The size of the object created to be created in capture in bytes + */ + default long getUsedMemory() + { + throw new UnsupportedOperationException(); + } }