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 2d0e35f..ce03dd6 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 @@ -1,5 +1,5 @@ /* - * Copyright 2017 the original author or authors. + * Copyright 2017-2018 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -44,7 +44,7 @@ import org.springframework.core.convert.converter.Converter; import org.springframework.core.serializer.support.DeserializingConverter; import org.springframework.integration.aws.support.AwsHeaders; import org.springframework.integration.endpoint.MessageProducerSupport; -import org.springframework.integration.metadata.MetadataStore; +import org.springframework.integration.metadata.ConcurrentMetadataStore; import org.springframework.integration.metadata.SimpleMetadataStore; import org.springframework.integration.support.AbstractIntegrationMessageBuilder; import org.springframework.integration.support.ErrorMessageStrategy; @@ -100,7 +100,7 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i private String consumerGroup = "SpringIntegration"; - private MetadataStore checkpointStore = new SimpleMetadataStore(); + private ConcurrentMetadataStore checkpointStore = new SimpleMetadataStore(); private Executor dispatcherExecutor; @@ -167,7 +167,7 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i this.consumerGroup = consumerGroup; } - public void setCheckpointStore(MetadataStore checkpointStore) { + public void setCheckpointStore(ConcurrentMetadataStore checkpointStore) { Assert.notNull(checkpointStore, "'checkpointStore' must not be null"); this.checkpointStore = checkpointStore; } diff --git a/src/main/java/org/springframework/integration/aws/inbound/kinesis/ShardCheckpointer.java b/src/main/java/org/springframework/integration/aws/inbound/kinesis/ShardCheckpointer.java index ff8bb5f..9d43c81 100644 --- a/src/main/java/org/springframework/integration/aws/inbound/kinesis/ShardCheckpointer.java +++ b/src/main/java/org/springframework/integration/aws/inbound/kinesis/ShardCheckpointer.java @@ -24,6 +24,7 @@ import java.util.List; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import org.springframework.integration.metadata.ConcurrentMetadataStore; import org.springframework.integration.metadata.MetadataStore; import com.amazonaws.services.kinesis.model.Record; @@ -43,7 +44,7 @@ class ShardCheckpointer implements Checkpointer { private static final Log logger = LogFactory.getLog(ShardCheckpointer.class); - private final MetadataStore checkpointStore; + private final ConcurrentMetadataStore checkpointStore; private final String key; @@ -51,7 +52,7 @@ class ShardCheckpointer implements Checkpointer { private volatile boolean active = true; - ShardCheckpointer(MetadataStore checkpointStore, String key) { + ShardCheckpointer(ConcurrentMetadataStore checkpointStore, String key) { this.checkpointStore = checkpointStore; this.key = key; } @@ -64,11 +65,15 @@ class ShardCheckpointer implements Checkpointer { @Override public boolean checkpoint(String sequenceNumber) { if (this.active) { - String existingSequence = this.checkpointStore.get(this.key); + String existingSequence = getCheckpoint(); if (existingSequence == null || new BigInteger(existingSequence).compareTo(new BigInteger(sequenceNumber)) < 0) { - this.checkpointStore.put(this.key, sequenceNumber); - return true; + if (existingSequence != null) { + return this.checkpointStore.replace(this.key, existingSequence, sequenceNumber); + } + else { + return this.checkpointStore.putIfAbsent(this.key, sequenceNumber) == null; + } } } else { 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 4412887..09c0714 100644 --- a/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java +++ b/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2017 the original author or authors. + * Copyright 2017-2018 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -46,6 +46,7 @@ import org.springframework.integration.aws.inbound.kinesis.ListenerMode; import org.springframework.integration.aws.support.AwsHeaders; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.config.EnableIntegration; +import org.springframework.integration.metadata.ConcurrentMetadataStore; import org.springframework.integration.metadata.MetadataStore; import org.springframework.integration.metadata.SimpleMetadataStore; import org.springframework.integration.test.util.TestUtils; @@ -297,7 +298,7 @@ public class KinesisMessageDrivenChannelAdapterTests { } @Bean - public MetadataStore checkpointStore() { + public ConcurrentMetadataStore checkpointStore() { SimpleMetadataStore simpleMetadataStore = new SimpleMetadataStore(); String testKey = "SpringIntegration" + ":" + STREAM1 + ":" + "3"; simpleMetadataStore.put(testKey, "1"); diff --git a/src/test/java/org/springframework/integration/aws/kinesis/KinesisIntegrationTests.java b/src/test/java/org/springframework/integration/aws/kinesis/KinesisIntegrationTests.java index 2989c4c..89a894a 100644 --- a/src/test/java/org/springframework/integration/aws/kinesis/KinesisIntegrationTests.java +++ b/src/test/java/org/springframework/integration/aws/kinesis/KinesisIntegrationTests.java @@ -19,6 +19,8 @@ package org.springframework.integration.aws.kinesis; import static org.assertj.core.api.Assertions.assertThat; import java.util.Date; +import java.util.HashSet; +import java.util.Set; import org.junit.AfterClass; import org.junit.BeforeClass; @@ -37,6 +39,8 @@ import org.springframework.integration.aws.outbound.KinesisMessageHandler; import org.springframework.integration.aws.support.AwsHeaders; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.config.EnableIntegration; +import org.springframework.integration.metadata.ConcurrentMetadataStore; +import org.springframework.integration.metadata.SimpleMetadataStore; import org.springframework.integration.support.MessageBuilder; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; @@ -81,7 +85,7 @@ public class KinesisIntegrationTests { } @Test - public void testKinesisInboundOutbound() { + public void testKinesisInboundOutbound() throws InterruptedException { this.kinesisSendChannel.send( MessageBuilder.withPayload("foo") .setHeader(AwsHeaders.STREAM, TEST_STREAM) @@ -103,6 +107,29 @@ public class KinesisIntegrationTests { assertThat(((Exception) errorMessage.getPayload()).getMessage()) .contains("Channel 'kinesisReceiveChannel' expected one of the following datataypes " + "[class java.util.Date], but received [class java.lang.String]"); + + + for (int i = 0; i < 1000; i++) { + this.kinesisSendChannel.send( + MessageBuilder.withPayload(new Date()) + .setHeader(AwsHeaders.STREAM, TEST_STREAM) + .build()); + } + + Set receivedSequences = new HashSet<>(); + + + for (int i = 0; i < 1000; i++) { + receive = this.kinesisReceiveChannel.receive(10_000); + assertThat(receive).isNotNull(); + String sequenceNumber = receive.getHeaders().get(AwsHeaders.RECEIVED_SEQUENCE_NUMBER, String.class); + assertThat(receivedSequences.add(sequenceNumber)).isTrue(); + } + + assertThat(receivedSequences.size()).isEqualTo(1000); + + receive = this.kinesisReceiveChannel.receive(10); + assertThat(receive).isNull(); } @Configuration @@ -118,15 +145,40 @@ public class KinesisIntegrationTests { } @Bean - public KinesisMessageDrivenChannelAdapter kinesisInboundChannelChannel() { + public ConcurrentMetadataStore checkpointStore() { + return new SimpleMetadataStore(); + } + + private KinesisMessageDrivenChannelAdapter kinesisMessageDrivenChannelAdapter() { KinesisMessageDrivenChannelAdapter adapter = new KinesisMessageDrivenChannelAdapter(KINESIS_LOCAL_RUNNING.getKinesis(), TEST_STREAM); adapter.setOutputChannel(kinesisReceiveChannel()); adapter.setErrorChannel(errorChannel()); adapter.setErrorMessageStrategy(new KinesisMessageHeaderErrorMessageStrategy()); + adapter.setCheckpointStore(checkpointStore()); return adapter; } + @Bean + public KinesisMessageDrivenChannelAdapter kinesisInboundChannelChannel1() { + return kinesisMessageDrivenChannelAdapter(); + } + + @Bean + public KinesisMessageDrivenChannelAdapter kinesisInboundChannelChannel2() { + return kinesisMessageDrivenChannelAdapter(); + } + + @Bean + public KinesisMessageDrivenChannelAdapter kinesisInboundChannelChannel3() { + return kinesisMessageDrivenChannelAdapter(); + } + + @Bean + public KinesisMessageDrivenChannelAdapter kinesisInboundChannelChannel4() { + return kinesisMessageDrivenChannelAdapter(); + } + @Bean public PollableChannel kinesisReceiveChannel() { QueueChannel queueChannel = new QueueChannel();