This is an automated email from the ASF dual-hosted git repository.

dannycranmer pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/flink-connector-aws.git


The following commit(s) were added to refs/heads/main by this push:
     new 8d29147  [FLINK-35115][Connectors/Kinesis] Allow kinesis consumer to 
snapshotState after cancelling operator
8d29147 is described below

commit 8d29147b9e6c0a7d27399662c6023ad634363764
Author: Aleksandr Pilipenko <[email protected]>
AuthorDate: Wed Apr 17 17:05:23 2024 +0100

    [FLINK-35115][Connectors/Kinesis] Allow kinesis consumer to snapshotState 
after cancelling operator
---
 .../connectors/kinesis/FlinkKinesisConsumer.java   |  13 +-
 .../kinesis/FlinkKinesisConsumerTest.java          | 304 +++++++++++----------
 2 files changed, 179 insertions(+), 138 deletions(-)

diff --git 
a/flink-connector-aws/flink-connector-kinesis/src/main/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumer.java
 
b/flink-connector-aws/flink-connector-kinesis/src/main/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumer.java
index c229a1c..06d0acc 100644
--- 
a/flink-connector-aws/flink-connector-kinesis/src/main/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumer.java
+++ 
b/flink-connector-aws/flink-connector-kinesis/src/main/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumer.java
@@ -148,7 +148,17 @@ public class FlinkKinesisConsumer<T> extends 
RichParallelSourceFunction<T>
     private transient HashMap<StreamShardMetadata.EquivalenceWrapper, 
SequenceNumber>
             sequenceNumsToRestore;
 
+    /**
+     * Flag used to control reading from Kinesis: source will read data while 
value is true. Changed
+     * to false after {@link #cancel()} has been called.
+     */
     private volatile boolean running = true;
+    /**
+     * Flag identifying that operator had been closed. True only after {@link 
#close()} has been
+     * called. Used to control behaviour of snapshotState: state can be 
persisted after operator has
+     * been cancelled (during stop-with-savepoint workflow), but not after 
operator has been closed.
+     */
+    private volatile boolean closed = false;
 
     // ------------------------------------------------------------------------
     //  State for Checkpoint
@@ -419,6 +429,7 @@ public class FlinkKinesisConsumer<T> extends 
RichParallelSourceFunction<T>
     @Override
     public void close() throws Exception {
         cancel();
+        closed = true;
         // safe-guard when the fetcher has been interrupted, make sure to not 
leak resources
         // application might be stopped before connector subtask has been 
started
         // so we must check if the fetcher is actually created
@@ -478,7 +489,7 @@ public class FlinkKinesisConsumer<T> extends 
RichParallelSourceFunction<T>
 
     @Override
     public void snapshotState(FunctionSnapshotContext context) throws 
Exception {
-        if (!running) {
+        if (closed) {
             LOG.debug("snapshotState() called on closed source; returning 
null.");
         } else {
             if (LOG.isDebugEnabled()) {
diff --git 
a/flink-connector-aws/flink-connector-kinesis/src/test/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumerTest.java
 
b/flink-connector-aws/flink-connector-kinesis/src/test/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumerTest.java
index f8cd5ab..82836df 100644
--- 
a/flink-connector-aws/flink-connector-kinesis/src/test/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumerTest.java
+++ 
b/flink-connector-aws/flink-connector-kinesis/src/test/java/org/apache/flink/streaming/connectors/kinesis/FlinkKinesisConsumerTest.java
@@ -39,6 +39,7 @@ import org.apache.flink.streaming.api.TimeCharacteristic;
 import org.apache.flink.streaming.api.functions.source.SourceFunction;
 import 
org.apache.flink.streaming.api.functions.timestamps.BoundedOutOfOrdernessTimestampExtractor;
 import org.apache.flink.streaming.api.operators.StreamSource;
+import 
org.apache.flink.streaming.api.operators.collect.utils.MockFunctionSnapshotContext;
 import org.apache.flink.streaming.api.watermark.Watermark;
 import org.apache.flink.streaming.api.windowing.time.Time;
 import 
org.apache.flink.streaming.connectors.kinesis.config.ConsumerConfigConstants;
@@ -71,7 +72,6 @@ import 
com.amazonaws.services.kinesis.model.SequenceNumberRange;
 import com.amazonaws.services.kinesis.model.Shard;
 import org.junit.Test;
 import org.junit.runner.RunWith;
-import org.mockito.Matchers;
 import org.mockito.MockedStatic;
 import org.mockito.Mockito;
 import org.powermock.api.mockito.PowerMockito;
@@ -97,10 +97,13 @@ import java.util.concurrent.atomic.AtomicReference;
 import java.util.function.Supplier;
 
 import static org.assertj.core.api.Assertions.assertThat;
+import static org.mockito.ArgumentMatchers.any;
 import static org.mockito.Mockito.mock;
 import static org.mockito.Mockito.mockStatic;
 import static org.mockito.Mockito.never;
 import static org.mockito.Mockito.spy;
+import static org.mockito.Mockito.times;
+import static org.mockito.Mockito.verify;
 import static org.mockito.Mockito.when;
 
 /** Suite of FlinkKinesisConsumer tests for the methods called throughout the 
source life cycle. */
@@ -116,53 +119,16 @@ public class FlinkKinesisConsumerTest extends TestLogger {
     public void testUseRestoredStateForSnapshotIfFetcherNotInitialized() 
throws Exception {
         Properties config = TestUtils.getStandardProperties();
 
-        List<Tuple2<StreamShardMetadata, SequenceNumber>> globalUnionState = 
new ArrayList<>(4);
-        globalUnionState.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(0)))),
-                        new SequenceNumber("1")));
-        globalUnionState.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(1)))),
-                        new SequenceNumber("1")));
-        globalUnionState.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(2)))),
-                        new SequenceNumber("1")));
-        globalUnionState.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(3)))),
-                        new SequenceNumber("1")));
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> globalUnionState =
+                Arrays.asList(
+                        createShardState("fakeStream", 0, "1"),
+                        createShardState("fakeStream", 1, "1"),
+                        createShardState("fakeStream", 2, "1"),
+                        createShardState("fakeStream", 3, "1"));
 
         TestingListState<Tuple2<StreamShardMetadata, SequenceNumber>> 
listState =
                 new TestingListState<>();
-        for (Tuple2<StreamShardMetadata, SequenceNumber> state : 
globalUnionState) {
-            listState.add(state);
-        }
+        listState.addAll(globalUnionState);
 
         FlinkKinesisConsumer<String> consumer =
                 new FlinkKinesisConsumer<>("fakeStream", new 
SimpleStringSchema(), config);
@@ -170,7 +136,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         consumer.setRuntimeContext(context);
 
         OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
                 .thenReturn(listState);
 
         StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
@@ -199,52 +165,14 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         // 
----------------------------------------------------------------------
         // setup config, initial state and expected state snapshot
         // 
----------------------------------------------------------------------
-        Properties config = TestUtils.getStandardProperties();
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> initialState =
+                Collections.singletonList(createShardState("fakeStream", 0, 
"1"));
 
-        ArrayList<Tuple2<StreamShardMetadata, SequenceNumber>> initialState = 
new ArrayList<>(1);
-        initialState.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream1",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(0)))),
-                        new SequenceNumber("1")));
-
-        ArrayList<Tuple2<StreamShardMetadata, SequenceNumber>> 
expectedStateSnapshot =
-                new ArrayList<>(3);
-        expectedStateSnapshot.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream1",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(0)))),
-                        new SequenceNumber("12")));
-        expectedStateSnapshot.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream1",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(1)))),
-                        new SequenceNumber("11")));
-        expectedStateSnapshot.add(
-                Tuple2.of(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream1",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(2)))),
-                        new SequenceNumber("31")));
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> 
expectedStateSnapshot =
+                Arrays.asList(
+                        createShardState("fakeStream", 0, "12"),
+                        createShardState("fakeStream", 1, "11"),
+                        createShardState("fakeStream", 2, "31"));
 
         // 
----------------------------------------------------------------------
         // mock operator state backend and initial state for initializeState()
@@ -252,17 +180,59 @@ public class FlinkKinesisConsumerTest extends TestLogger {
 
         TestingListState<Tuple2<StreamShardMetadata, SequenceNumber>> 
listState =
                 new TestingListState<>();
-        for (Tuple2<StreamShardMetadata, SequenceNumber> state : initialState) 
{
-            listState.add(state);
+        listState.addAll(initialState);
+
+        // 
----------------------------------------------------------------------
+        // mock a running fetcher and its state for snapshot
+        // 
----------------------------------------------------------------------
+
+        HashMap<StreamShardMetadata, SequenceNumber> stateSnapshot = new 
HashMap<>();
+        for (Tuple2<StreamShardMetadata, SequenceNumber> tuple : 
expectedStateSnapshot) {
+            stateSnapshot.put(tuple.f0, tuple.f1);
         }
 
-        OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
-                .thenReturn(listState);
+        KinesisDataFetcher<String> mockedFetcher = 
mock(KinesisDataFetcher.class);
+        when(mockedFetcher.snapshotState()).thenReturn(stateSnapshot);
 
-        StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
-        
when(initializationContext.getOperatorStateStore()).thenReturn(operatorStateStore);
-        when(initializationContext.isRestored()).thenReturn(true);
+        // 
----------------------------------------------------------------------
+        // create a consumer and test the snapshotState()
+        // 
----------------------------------------------------------------------
+
+        FlinkKinesisConsumer<String> mockedConsumer =
+                prepareMockedConsumer(
+                        "fakeStream", new SimpleStringSchema(), mockedFetcher, 
listState);
+
+        mockedConsumer.snapshotState(mock(FunctionSnapshotContext.class));
+
+        // 
----------------------------------------------------------------------
+        // verify that state had been updated
+        // 
----------------------------------------------------------------------
+        assertThat(listState.clearCalled).isTrue();
+        assertThat(listState.getList())
+                .hasSize(3)
+                .doesNotContainAnyElementsOf(initialState)
+                .containsAll(expectedStateSnapshot);
+    }
+
+    @Test
+    public void testSnapshotStateChangesAfterCancel() throws Exception {
+
+        // 
----------------------------------------------------------------------
+        // setup config, initial state and expected state snapshot
+        // 
----------------------------------------------------------------------
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> initialState =
+                Collections.singletonList(createShardState("fakeStream", 0, 
"11"));
+
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> 
expectedStateSnapshot =
+                Collections.singletonList(createShardState("fakeStream", 0, 
"12"));
+
+        // 
----------------------------------------------------------------------
+        // mock operator state backend and initial state for initializeState()
+        // 
----------------------------------------------------------------------
+
+        TestingListState<Tuple2<StreamShardMetadata, SequenceNumber>> 
listState =
+                new TestingListState<>();
+        listState.addAll(initialState);
 
         // 
----------------------------------------------------------------------
         // mock a running fetcher and its state for snapshot
@@ -273,43 +243,99 @@ public class FlinkKinesisConsumerTest extends TestLogger {
             stateSnapshot.put(tuple.f0, tuple.f1);
         }
 
-        KinesisDataFetcher mockedFetcher = mock(KinesisDataFetcher.class);
+        KinesisDataFetcher<String> mockedFetcher = 
mock(KinesisDataFetcher.class);
         when(mockedFetcher.snapshotState()).thenReturn(stateSnapshot);
 
         // 
----------------------------------------------------------------------
         // create a consumer and test the snapshotState()
         // 
----------------------------------------------------------------------
 
-        FlinkKinesisConsumer<String> consumer =
-                new FlinkKinesisConsumer<>("fakeStream", new 
SimpleStringSchema(), config);
-        FlinkKinesisConsumer<?> mockedConsumer = spy(consumer);
+        FlinkKinesisConsumer<String> mockedConsumer =
+                prepareMockedConsumer(
+                        "fakeStream", new SimpleStringSchema(), mockedFetcher, 
listState);
+
+        mockedConsumer.cancel();
+        mockedConsumer.snapshotState(new MockFunctionSnapshotContext(2));
+        verify(mockedFetcher, times(1)).snapshotState();
+
+        // 
----------------------------------------------------------------------
+        // verify that state had been updated
+        // 
----------------------------------------------------------------------
+        assertThat(listState.isClearCalled()).isTrue();
+        assertThat(listState.getList())
+                .hasSize(1)
+                .doesNotContainAnyElementsOf(initialState)
+                .containsAll(expectedStateSnapshot);
+    }
+
+    @Test
+    public void testSnapshotStateNotChangedAfterClose() throws Exception {
+
+        // 
----------------------------------------------------------------------
+        // setup initial state
+        // 
----------------------------------------------------------------------
+
+        List<Tuple2<StreamShardMetadata, SequenceNumber>> initialState =
+                Collections.singletonList(createShardState("fakeStream", 0, 
"11"));
+
+        // 
----------------------------------------------------------------------
+        // mock initial state
+        // 
----------------------------------------------------------------------
+
+        TestingListState<Tuple2<StreamShardMetadata, SequenceNumber>> 
listState =
+                new TestingListState<>();
+        listState.addAll(initialState);
+
+        // 
----------------------------------------------------------------------
+        // mock a running fetcher and its state for snapshot
+        // 
----------------------------------------------------------------------
+
+        KinesisDataFetcher<String> mockedFetcher = 
mock(KinesisDataFetcher.class);
+        when(mockedFetcher.snapshotState()).thenReturn(new HashMap<>());
+
+        // 
----------------------------------------------------------------------
+        // create a consumer and test the snapshotState()
+        // 
----------------------------------------------------------------------
+
+        FlinkKinesisConsumer<String> mockedConsumer =
+                prepareMockedConsumer(
+                        "fakeStream", new SimpleStringSchema(), mockedFetcher, 
listState);
+
+        mockedConsumer.close();
+        mockedConsumer.snapshotState(new MockFunctionSnapshotContext(3));
+        verify(mockedFetcher, never()).snapshotState();
+        assertThat(listState.isClearCalled()).isFalse();
+        assertThat(listState.getList()).containsAll(initialState);
+    }
+
+    private <T> FlinkKinesisConsumer<T> prepareMockedConsumer(
+            String streamName,
+            DeserializationSchema<T> schema,
+            KinesisDataFetcher<T> fetcher,
+            ListState<?> listState)
+            throws Exception {
+
+        Properties config = TestUtils.getStandardProperties();
+
+        OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
+                .thenReturn(listState);
+
+        StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
+        
when(initializationContext.getOperatorStateStore()).thenReturn(operatorStateStore);
+        when(initializationContext.isRestored()).thenReturn(true);
+
+        FlinkKinesisConsumer<T> consumer = new 
FlinkKinesisConsumer<>(streamName, schema, config);
+        FlinkKinesisConsumer<T> mockedConsumer = spy(consumer);
 
         RuntimeContext context = new MockStreamingRuntimeContext(true, 1, 0);
 
         mockedConsumer.setRuntimeContext(context);
         mockedConsumer.initializeState(initializationContext);
         mockedConsumer.open(new Configuration());
-        Whitebox.setInternalState(
-                mockedConsumer, "fetcher", mockedFetcher); // mock consumer as 
running.
-
-        mockedConsumer.snapshotState(mock(FunctionSnapshotContext.class));
-
-        assertThat(listState.clearCalled).isTrue();
-        assertThat(listState.getList()).hasSize(3);
-
-        for (Tuple2<StreamShardMetadata, SequenceNumber> state : initialState) 
{
-            for (Tuple2<StreamShardMetadata, SequenceNumber> currentState : 
listState.getList()) {
-                assertThat(currentState).isNotEqualTo(state);
-            }
-        }
+        Whitebox.setInternalState(mockedConsumer, "fetcher", fetcher); // mock 
consumer as running.
 
-        for (Tuple2<StreamShardMetadata, SequenceNumber> state : 
expectedStateSnapshot) {
-            boolean hasOneIsSame = false;
-            for (Tuple2<StreamShardMetadata, SequenceNumber> currentState : 
listState.getList()) {
-                hasOneIsSame = hasOneIsSame || state.equals(currentState);
-            }
-            assertThat(hasOneIsSame).isTrue();
-        }
+        return mockedConsumer;
     }
 
     /**
@@ -322,16 +348,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
     public void testExplicitStateSerializerCompatibility() throws Exception {
         ExecutionConfig executionConfig = new ExecutionConfig();
 
-        Tuple2<StreamShardMetadata, SequenceNumber> tuple =
-                new Tuple2<>(
-                        KinesisDataFetcher.convertToStreamShardMetadata(
-                                new StreamShardHandle(
-                                        "fakeStream",
-                                        new Shard()
-                                                .withShardId(
-                                                        KinesisShardIdGenerator
-                                                                
.generateFromShardOrder(0)))),
-                        new SequenceNumber("1"));
+        Tuple2<StreamShardMetadata, SequenceNumber> tuple = 
createShardState("fakeStream", 0, "1");
 
         // This is how serializer was created implicitly using a 
TypeInformation
         // and since SequenceNumber is GenericType, Flink falls back to Kryo
@@ -358,6 +375,19 @@ public class FlinkKinesisConsumerTest extends TestLogger {
                 .isEqualTo(actualTuple);
     }
 
+    private Tuple2<StreamShardMetadata, SequenceNumber> createShardState(
+            String streamName, int shardNumber, String sequenceNumber) {
+        return Tuple2.of(
+                KinesisDataFetcher.convertToStreamShardMetadata(
+                        new StreamShardHandle(
+                                streamName,
+                                new Shard()
+                                        .withShardId(
+                                                
KinesisShardIdGenerator.generateFromShardOrder(
+                                                        shardNumber)))),
+                new SequenceNumber(sequenceNumber));
+    }
+
     // ----------------------------------------------------------------------
     // Tests related to fetcher initialization
     // ----------------------------------------------------------------------
@@ -401,7 +431,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         }
 
         OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
                 .thenReturn(listState);
 
         StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
@@ -478,7 +508,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         }
 
         OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
                 .thenReturn(listState);
 
         StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
@@ -586,7 +616,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         }
 
         OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
                 .thenReturn(listState);
 
         StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);
@@ -720,7 +750,7 @@ public class FlinkKinesisConsumerTest extends TestLogger {
         }
 
         OperatorStateStore operatorStateStore = mock(OperatorStateStore.class);
-        
when(operatorStateStore.getUnionListState(Matchers.any(ListStateDescriptor.class)))
+        
when(operatorStateStore.getUnionListState(any(ListStateDescriptor.class)))
                 .thenReturn(listState);
 
         StateInitializationContext initializationContext = 
mock(StateInitializationContext.class);

Reply via email to