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..81f8df553119 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 @@ -21,6 +21,8 @@ import java.util.Collections; import java.util.HashSet; import java.util.Iterator; +import java.util.List; +import java.util.ArrayList; import java.util.Map; import java.util.Objects; import java.util.Set; @@ -620,18 +622,18 @@ private static class FlinkOrderedListState implements OrderedListState { @Override public Iterable> readRange(Instant minTimestamp, Instant limitTimestamp) { - return readAsMap().subMap(minTimestamp, limitTimestamp).values(); + return Iterables.concat(readAsMap().subMap(minTimestamp, limitTimestamp).values()); } @Override public void clearRange(Instant minTimestamp, Instant limitTimestamp) { - SortedMap> sortedMap = readAsMap(); + SortedMap>> sortedMap = readAsMap(); sortedMap.subMap(minTimestamp, limitTimestamp).clear(); try { ListState> partitionedState = flinkStateBackend.getPartitionedState( namespace, namespaceSerializer, flinkStateDescriptor); - partitionedState.update(Lists.newArrayList(sortedMap.values())); + partitionedState.update(Lists.newArrayList(Iterables.concat(sortedMap.values()))); } catch (Exception e) { throw new RuntimeException("Error adding to bag state.", e); } @@ -680,10 +682,10 @@ public ReadableState readLater() { @Override @Nullable public Iterable> read() { - return readAsMap().values(); + return Iterables.concat(readAsMap().values()); } - private SortedMap> readAsMap() { + private SortedMap>> readAsMap() { Iterable> listValues; try { ListState> partitionedState = @@ -694,9 +696,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; }