From b5f238bd789a193c3d6959e581d793f6445f5c46 Mon Sep 17 00:00:00 2001 From: Dirk Bonhomme Date: Thu, 18 Jul 2019 17:35:25 +0200 Subject: [PATCH] GH-111: Implement batch mode for Kcl adapter Fixes https://github.com/spring-projects/spring-integration-aws/issues/111 --- .../KclMessageDrivenChannelAdapter.java | 161 +++++++++++++----- .../KinesisMessageDrivenChannelAdapter.java | 113 +++++++----- ...nesisMessageDrivenChannelAdapterTests.java | 2 +- 3 files changed, 188 insertions(+), 88 deletions(-) diff --git a/src/main/java/org/springframework/integration/aws/inbound/kinesis/KclMessageDrivenChannelAdapter.java b/src/main/java/org/springframework/integration/aws/inbound/kinesis/KclMessageDrivenChannelAdapter.java index a00b6b6..59f1b31 100644 --- a/src/main/java/org/springframework/integration/aws/inbound/kinesis/KclMessageDrivenChannelAdapter.java +++ b/src/main/java/org/springframework/integration/aws/inbound/kinesis/KclMessageDrivenChannelAdapter.java @@ -16,11 +16,17 @@ package org.springframework.integration.aws.inbound.kinesis; +import java.util.ArrayList; import java.util.Arrays; import java.util.List; import java.util.UUID; +import java.util.stream.Collectors; + +import javax.annotation.Nullable; import org.springframework.core.AttributeAccessor; +import org.springframework.core.convert.converter.Converter; +import org.springframework.core.serializer.support.DeserializingConverter; import org.springframework.core.task.SimpleAsyncTaskExecutor; import org.springframework.core.task.TaskExecutor; import org.springframework.core.task.support.ExecutorServiceAdapter; @@ -101,6 +107,10 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { private int consumerBackoff; + private Converter converter = new DeserializingConverter(); + + private ListenerMode listenerMode = ListenerMode.record; + private long checkpointsInterval = 5_000L; private CheckpointMode checkpointMode = CheckpointMode.batch; @@ -167,6 +177,20 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { this.consumerBackoff = Math.max(1000, consumerBackoff); } + /** + * Specify a {@link Converter} to deserialize the {@code byte[]} from record's body. + * Can be {@code null} meaning no deserialization. + * @param converter the {@link Converter} to use or null + */ + public void setConverter(Converter converter) { + this.converter = converter; + } + + public void setListenerMode(ListenerMode listenerMode) { + Assert.notNull(listenerMode, "'listenerMode' must not be null"); + this.listenerMode = listenerMode; + } + /** * Sets the interval between 2 checkpoints. * @param checkpointsInterval interval between 2 checkpoints (in milliseconds) @@ -226,6 +250,13 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { @Override protected void doStart() { super.doStart(); + + if (ListenerMode.batch.equals(this.listenerMode) && CheckpointMode.record.equals(this.checkpointMode)) { + this.checkpointMode = CheckpointMode.batch; + logger.warn("The 'checkpointMode' is overridden from [CheckpointMode.record] to [CheckpointMode.batch] " + + "because it does not make sense in case of [ListenerMode.batch]."); + } + this.executor.execute(this.scheduler); } @@ -288,47 +319,63 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { if (logger.isDebugEnabled()) { logger.debug("Processing " + records.size() + " records from " + this.shardId); } - for (Record record : records) { - try { - processSingleRecord(record, checkpointer); - } - catch (Throwable t) { - logger.warn("Caught throwable while processing record " + record, t); - } - finally { - attributesHolder.remove(); - // Checkpoint once every checkpoint interval. - if (CheckpointMode.periodic.equals(KclMessageDrivenChannelAdapter.this.checkpointMode) && - System.currentTimeMillis() > nextCheckpointTimeInMillis) { - checkpoint(checkpointer); - this.nextCheckpointTimeInMillis = System.currentTimeMillis() + checkpointsInterval; + + try { + if (ListenerMode.record.equals(KclMessageDrivenChannelAdapter.this.listenerMode)) { + for (Record record : records) { + processSingleRecord(record, checkpointer); + checkpointIfRecordMode(checkpointer, record); + checkpointIfPeriodicMode(checkpointer, record); } } + else if (ListenerMode.batch.equals(KclMessageDrivenChannelAdapter.this.listenerMode)) { + processMultipleRecords(records, checkpointer); + checkpointIfPeriodicMode(checkpointer, null); + } + checkpointIfBatchMode(checkpointer); } - - // checkpoint if needed - if (CheckpointMode.batch.equals(KclMessageDrivenChannelAdapter.this.checkpointMode)) { - checkpoint(checkpointer); + finally { + attributesHolder.remove(); } } - /** - * Process a single record. - * @param record The record to be processed. - * @param checkpointer the checkpointer to use if the checkpointMode is record - */ private void processSingleRecord(Record record, IRecordProcessorCheckpointer checkpointer) { - // Convert AWS Record in Spring Message. - performSend(prepareMessageForRecord(record, checkpointer), record); - - // checkpoint if needed - if (CheckpointMode.record.equals(KclMessageDrivenChannelAdapter.this.checkpointMode)) { - checkpoint(checkpointer); - } + performSend(prepareMessageForRecord(record), record, checkpointer); } - private AbstractIntegrationMessageBuilder prepareMessageForRecord(Record record, - IRecordProcessorCheckpointer checkpointer) { + private void processMultipleRecords(List records, IRecordProcessorCheckpointer checkpointer) { + Object payload = records; + + if (KclMessageDrivenChannelAdapter.this.embeddedHeadersMapper != null) { + payload = records.stream().map(this::prepareMessageForRecord).collect(Collectors.toList()); + } + + final List partitionKeys; + final List sequenceNumbers; + if (KclMessageDrivenChannelAdapter.this.converter != null) { + partitionKeys = new ArrayList<>(); + sequenceNumbers = new ArrayList<>(); + + payload = records.stream().map(r -> { + partitionKeys.add(r.getPartitionKey()); + sequenceNumbers.add(r.getSequenceNumber()); + + return KclMessageDrivenChannelAdapter.this.converter.convert(r.getData().array()); + }).collect(Collectors.toList()); + } + else { + partitionKeys = null; + sequenceNumbers = null; + } + + AbstractIntegrationMessageBuilder messageBuilder = getMessageBuilderFactory().withPayload(payload) + .setHeader(AwsHeaders.RECEIVED_PARTITION_KEY, partitionKeys) + .setHeader(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, sequenceNumbers); + + performSend(messageBuilder, records, checkpointer); + } + + private AbstractIntegrationMessageBuilder prepareMessageForRecord(Record record) { Object payload = record.getData().array(); Message messageToUse = null; @@ -347,11 +394,13 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { } } + if (payload instanceof byte[] && KclMessageDrivenChannelAdapter.this.converter != null) { + payload = KclMessageDrivenChannelAdapter.this.converter.convert((byte[]) payload); + } + AbstractIntegrationMessageBuilder messageBuilder = getMessageBuilderFactory().withPayload(payload) .setHeader(AwsHeaders.RECEIVED_PARTITION_KEY, record.getPartitionKey()) - .setHeader(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, record.getSequenceNumber()) - .setHeader(AwsHeaders.RECEIVED_STREAM, KclMessageDrivenChannelAdapter.this.stream) - .setHeader(AwsHeaders.SHARD, this.shardId); + .setHeader(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, record.getSequenceNumber()); if (KclMessageDrivenChannelAdapter.this.bindSourceRecord) { messageBuilder.setHeader(IntegrationMessageHeaderAccessor.SOURCE_DATA, record); @@ -361,14 +410,18 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { messageBuilder.copyHeadersIfAbsent(messageToUse.getHeaders()); } + return messageBuilder; + } + + private void performSend(AbstractIntegrationMessageBuilder messageBuilder, Object rawRecord, + IRecordProcessorCheckpointer checkpointer) { + messageBuilder.setHeader(AwsHeaders.RECEIVED_STREAM, KclMessageDrivenChannelAdapter.this.stream) + .setHeader(AwsHeaders.SHARD, this.shardId); + if (CheckpointMode.manual.equals(KclMessageDrivenChannelAdapter.this.checkpointMode)) { messageBuilder.setHeader(AwsHeaders.CHECKPOINTER, checkpointer); } - return messageBuilder; - } - - private void performSend(AbstractIntegrationMessageBuilder messageBuilder, Object rawRecord) { Message messageToSend = messageBuilder.build(); setAttributesIfNecessary(rawRecord, messageToSend); try { @@ -397,13 +450,19 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { /** * Checkpoint with retries. * @param checkpointer checkpointer + * @param record last processed record */ - private void checkpoint(IRecordProcessorCheckpointer checkpointer) { + private void checkpoint(IRecordProcessorCheckpointer checkpointer, @Nullable Record record) { if (logger.isInfoEnabled()) { logger.info("Checkpointing shard " + shardId); } try { - checkpointer.checkpoint(); + if (record == null) { + checkpointer.checkpoint(); + } + else { + checkpointer.checkpoint(record); + } } catch (ShutdownException se) { // Ignore checkpoint if the processor instance has been shutdown (fail @@ -424,6 +483,26 @@ public class KclMessageDrivenChannelAdapter extends MessageProducerSupport { } } + private void checkpointIfBatchMode(IRecordProcessorCheckpointer checkpointer) { + if (CheckpointMode.batch.equals(KclMessageDrivenChannelAdapter.this.checkpointMode)) { + checkpoint(checkpointer, null); + } + } + + private void checkpointIfRecordMode(IRecordProcessorCheckpointer checkpointer, Record record) { + if (CheckpointMode.record.equals(KclMessageDrivenChannelAdapter.this.checkpointMode)) { + checkpoint(checkpointer, record); + } + } + + private void checkpointIfPeriodicMode(IRecordProcessorCheckpointer checkpointer, @Nullable Record record) { + if (CheckpointMode.periodic.equals(KclMessageDrivenChannelAdapter.this.checkpointMode) + && System.currentTimeMillis() > nextCheckpointTimeInMillis) { + checkpoint(checkpointer, record); + this.nextCheckpointTimeInMillis = System.currentTimeMillis() + checkpointsInterval; + } + } + @Override public void shutdown(IRecordProcessorCheckpointer checkpointer, ShutdownReason reason) { if (logger.isInfoEnabled()) { diff --git a/src/main/java/org/springframework/integration/aws/inbound/kinesis/KinesisMessageDrivenChannelAdapter.java b/src/main/java/org/springframework/integration/aws/inbound/kinesis/KinesisMessageDrivenChannelAdapter.java index 7a5c0ce..8064f1c 100644 --- a/src/main/java/org/springframework/integration/aws/inbound/kinesis/KinesisMessageDrivenChannelAdapter.java +++ b/src/main/java/org/springframework/integration/aws/inbound/kinesis/KinesisMessageDrivenChannelAdapter.java @@ -42,6 +42,8 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.locks.Lock; import java.util.stream.Collectors; +import javax.annotation.Nullable; + import org.springframework.beans.factory.DisposableBean; import org.springframework.core.AttributeAccessor; import org.springframework.core.convert.converter.Converter; @@ -86,6 +88,7 @@ import com.amazonaws.services.kinesis.model.StreamStatus; * @author Artem Bilan * @author Krzysztof Witkowski * @author Hervé Fortin + * @author Dirk Bonhomme * @since 1.1 */ @ManagedResource @@ -972,60 +975,54 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i this.checkpointer.setHighestSequence(records.get(records.size() - 1).getSequenceNumber()); - switch (KinesisMessageDrivenChannelAdapter.this.listenerMode) { - case record: + if (ListenerMode.record.equals(KinesisMessageDrivenChannelAdapter.this.listenerMode)) { for (Record record : records) { - performSend(prepareMessageForRecord(record), record); - - if (CheckpointMode.record.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode)) { - this.checkpointer.checkpoint(record.getSequenceNumber()); - } + processSingleRecord(record); + checkpointIfRecordMode(record); + checkpointIfPeriodicMode(record); } + } + else if (ListenerMode.batch.equals(KinesisMessageDrivenChannelAdapter.this.listenerMode)) { + processMultipleRecords(records); + checkpointIfPeriodicMode(null); + } + checkpointIfBatchMode(); + } - break; + private void processSingleRecord(Record record) { + performSend(prepareMessageForRecord(record), record); + } - case batch: - Object payload = records; + private void processMultipleRecords(List records) { + Object payload = records; - if (KinesisMessageDrivenChannelAdapter.this.embeddedHeadersMapper != null) { - payload = records.stream().map(this::prepareMessageForRecord).collect(Collectors.toList()); - } - - final List partitionKeys; - final List sequenceNumbers; - if (KinesisMessageDrivenChannelAdapter.this.converter != null) { - partitionKeys = new ArrayList<>(); - sequenceNumbers = new ArrayList<>(); - - payload = records.stream().map(r -> { - partitionKeys.add(r.getPartitionKey()); - sequenceNumbers.add(r.getSequenceNumber()); - - return KinesisMessageDrivenChannelAdapter.this.converter.convert(r.getData().array()); - }).collect(Collectors.toList()); - } - else { - partitionKeys = null; - sequenceNumbers = null; - } - - AbstractIntegrationMessageBuilder messageBuilder = getMessageBuilderFactory().withPayload(payload) - .setHeader(AwsHeaders.RECEIVED_PARTITION_KEY, partitionKeys) - .setHeader(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, sequenceNumbers); - - performSend(messageBuilder, records); - - break; + if (KinesisMessageDrivenChannelAdapter.this.embeddedHeadersMapper != null) { + payload = records.stream().map(this::prepareMessageForRecord).collect(Collectors.toList()); } - if (CheckpointMode.batch.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode)) { - this.checkpointer.checkpoint(); + final List partitionKeys; + final List sequenceNumbers; + if (KinesisMessageDrivenChannelAdapter.this.converter != null) { + partitionKeys = new ArrayList<>(); + sequenceNumbers = new ArrayList<>(); + + payload = records.stream().map(r -> { + partitionKeys.add(r.getPartitionKey()); + sequenceNumbers.add(r.getSequenceNumber()); + + return KinesisMessageDrivenChannelAdapter.this.converter.convert(r.getData().array()); + }).collect(Collectors.toList()); } - else if (CheckpointMode.periodic.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode) - && System.currentTimeMillis() > nextCheckpointTimeInMillis) { - this.checkpointer.checkpoint(); - this.nextCheckpointTimeInMillis = System.currentTimeMillis() + checkpointsInterval; + else { + partitionKeys = null; + sequenceNumbers = null; } + + AbstractIntegrationMessageBuilder messageBuilder = getMessageBuilderFactory().withPayload(payload) + .setHeader(AwsHeaders.RECEIVED_PARTITION_KEY, partitionKeys) + .setHeader(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, sequenceNumbers); + + performSend(messageBuilder, records); } private AbstractIntegrationMessageBuilder prepareMessageForRecord(Record record) { @@ -1045,7 +1042,6 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i } if (payload instanceof byte[] && KinesisMessageDrivenChannelAdapter.this.converter != null) { - payload = KinesisMessageDrivenChannelAdapter.this.converter.convert((byte[]) payload); } @@ -1083,6 +1079,31 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i } } + private void checkpointIfBatchMode() { + if (CheckpointMode.batch.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode)) { + this.checkpointer.checkpoint(); + } + } + + private void checkpointIfRecordMode(Record record) { + if (CheckpointMode.record.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode)) { + this.checkpointer.checkpoint(record.getSequenceNumber()); + } + } + + private void checkpointIfPeriodicMode(@Nullable Record record) { + if (CheckpointMode.periodic.equals(KinesisMessageDrivenChannelAdapter.this.checkpointMode) + && System.currentTimeMillis() > nextCheckpointTimeInMillis) { + if (record == null) { + this.checkpointer.checkpoint(); + } + else { + this.checkpointer.checkpoint(record.getSequenceNumber()); + } + this.nextCheckpointTimeInMillis = System.currentTimeMillis() + checkpointsInterval; + } + } + @Override public String toString() { return "ShardConsumer{" + "shardOffset=" + this.shardOffset + ", state=" + this.state + '}'; diff --git a/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java b/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java index b2d405b..c36841a 100644 --- a/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java +++ b/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java @@ -195,7 +195,7 @@ public class KinesisMessageDrivenChannelAdapterTests { @Test @SuppressWarnings("rawtypes") - public void testReshadring() throws InterruptedException { + public void testResharding() throws InterruptedException { this.reshardingChannelAdapter.start(); assertThat(this.kinesisChannel.receive(10000)).isNotNull();