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);