diff --git a/runners/flink/src/main/java/org/apache/beam/runners/flink/translation/wrappers/streaming/state/FlinkStateInternals.java b/runners/flink/src/main/java/org/apache/beam/runners/flink/translation/wrappers/streaming/state/FlinkStateInternals.java index de244c1c8a31..a5870504c377 100644 --- a/runners/flink/src/main/java/org/apache/beam/runners/flink/translation/wrappers/streaming/state/FlinkStateInternals.java +++ b/runners/flink/src/main/java/org/apache/beam/runners/flink/translation/wrappers/streaming/state/FlinkStateInternals.java @@ -20,8 +20,11 @@ import java.io.IOException; import java.util.Collections; import java.util.HashSet; +import java.util.ArrayList; import java.util.Iterator; +import java.util.List; import java.util.Map; +import java.util.stream.Collectors; import java.util.Objects; import java.util.Set; import java.util.SortedMap; @@ -620,18 +623,21 @@ private static class FlinkOrderedListState implements OrderedListState { @Override public Iterable> readRange(Instant minTimestamp, Instant limitTimestamp) { - return readAsMap().subMap(minTimestamp, limitTimestamp).values(); + return readAsMultimap().subMap(minTimestamp, limitTimestamp).values().stream() + .flatMap(List::stream) + .collect(Collectors.toList()); } @Override public void clearRange(Instant minTimestamp, Instant limitTimestamp) { - SortedMap> sortedMap = readAsMap(); + SortedMap>> sortedMap = readAsMultimap(); sortedMap.subMap(minTimestamp, limitTimestamp).clear(); try { ListState> partitionedState = flinkStateBackend.getPartitionedState( namespace, namespaceSerializer, flinkStateDescriptor); - partitionedState.update(Lists.newArrayList(sortedMap.values())); + partitionedState.update( + sortedMap.values().stream().flatMap(List::stream).collect(Collectors.toList())); } catch (Exception e) { throw new RuntimeException("Error adding to bag state.", e); } @@ -680,10 +686,12 @@ public ReadableState readLater() { @Override @Nullable public Iterable> read() { - return readAsMap().values(); + return readAsMultimap().values().stream() + .flatMap(List::stream) + .collect(Collectors.toList()); } - private SortedMap> readAsMap() { + private SortedMap>> readAsMultimap() { Iterable> listValues; try { ListState> partitionedState = @@ -694,9 +702,9 @@ private SortedMap> readAsMap() { throw new RuntimeException("Error reading state.", e); } - SortedMap> sortedMap = Maps.newTreeMap(); + SortedMap>> sortedMap = Maps.newTreeMap(); for (TimestampedValue value : listValues) { - sortedMap.put(value.getTimestamp(), value); + sortedMap.computeIfAbsent(value.getTimestamp(), k -> new ArrayList<>()).add(value); } return sortedMap; }