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 e01955d..235e491 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 @@ -35,10 +35,8 @@ import java.util.concurrent.CountDownLatch; import java.util.concurrent.Executor; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.Semaphore; import java.util.concurrent.TimeUnit; -import java.util.concurrent.locks.Condition; -import java.util.concurrent.locks.Lock; -import java.util.concurrent.locks.ReentrantLock; import org.springframework.beans.factory.DisposableBean; import org.springframework.core.convert.converter.Converter; @@ -119,7 +117,9 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i private CheckpointMode checkpointMode = CheckpointMode.batch; - private int recordsLimit = 25; + private int recordsLimit = 10000; + + private int idleBetweenPolls = 1000; private int consumerBackoff = 1000; @@ -201,9 +201,15 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i this.checkpointMode = checkpointMode; } + /** + * The maximum record to poll per on get-records request. + * Not greater then {@code 10000}. + * @param recordsLimit the number of records to for per on get-records request. + * @see GetRecordsRequest#setLimit + */ public void setRecordsLimit(int recordsLimit) { Assert.isTrue(recordsLimit > 0, "'recordsLimit' must be more than 0"); - this.recordsLimit = recordsLimit; + this.recordsLimit = Math.min(10000, recordsLimit); } public void setConsumerBackoff(int consumerBackoff) { @@ -237,6 +243,15 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i this.maxConcurrency = concurrency; } + /** + * The sleep interval in milliseconds used in the main loop between shards polling cycles. + * Defaults to {@code 1000}l minimum {@code 250}. + * @param idleBetweenPolls the interval to sleep between shards polling cycles. + */ + public void setIdleBetweenPolls(int idleBetweenPolls) { + this.idleBetweenPolls = Math.max(250, idleBetweenPolls); + } + @Override protected void onInit() { super.onInit(); @@ -649,19 +664,17 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i KinesisMessageDrivenChannelAdapter.this.shardOffsets.remove(shardOffset); } } - else { - try { - Thread.sleep(1); - } - catch (InterruptedException e) { - Thread.currentThread().interrupt(); - throw new IllegalStateException("ConsumerDispatcher Thread [" + this + - "] has been interrupted", e); - } - } } } } + + try { + Thread.sleep(KinesisMessageDrivenChannelAdapter.this.idleBetweenPolls); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("ConsumerDispatcher Thread [" + this + "] has been interrupted", e); + } } } @@ -696,6 +709,10 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i this.checkpointer = new ShardCheckpointer(KinesisMessageDrivenChannelAdapter.this.checkpointStore, key); } + void setNotifier(Runnable notifier) { + this.notifier = notifier; + } + void stop() { this.state = ConsumerState.STOP; if (this.notifier != null) { @@ -923,9 +940,6 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i '}'; } - public void setNotifier(Runnable notifier) { - this.notifier = notifier; - } } private enum ConsumerState { @@ -938,9 +952,7 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i private final Queue consumers = new ConcurrentLinkedQueue<>(); - private final Lock lock = new ReentrantLock(); - - private final Condition condition = this.lock.newCondition(); + private final Semaphore processBarrier = new Semaphore(0); private final Runnable notifier = new Runnable() { @@ -958,55 +970,44 @@ public class KinesisMessageDrivenChannelAdapter extends MessageProducerSupport i } void addConsumer(ShardConsumer shardConsumer) { - this.consumers.add(shardConsumer); shardConsumer.setNotifier(this.notifier); + this.consumers.add(shardConsumer); } void notifyBarrier() { - this.lock.lock(); - try { - this.condition.signalAll(); - } - finally { - this.lock.unlock(); - } + this.processBarrier.release(); } @Override public void run() { - this.lock.lock(); - try { - while (KinesisMessageDrivenChannelAdapter.this.active) { - try { - this.condition.await(); - } - catch (InterruptedException e) { - Thread.currentThread().interrupt(); + while (KinesisMessageDrivenChannelAdapter.this.active) { + try { + this.processBarrier.acquire(); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); - throw new IllegalStateException("ConsumerInvoker thread [" + this + "] has been interrupted", e); + throw new IllegalStateException("ConsumerInvoker thread [" + this + "] has been interrupted", e); + } + + for (Iterator iterator = this.consumers.iterator(); iterator.hasNext(); ) { + ShardConsumer shardConsumer = iterator.next(); + if (ConsumerState.STOP == shardConsumer.state) { + iterator.remove(); } - for (Iterator iterator = this.consumers.iterator(); iterator.hasNext(); ) { - ShardConsumer shardConsumer = iterator.next(); - if (ConsumerState.STOP == shardConsumer.state) { - iterator.remove(); - } - else { - if (shardConsumer.task != null) { - shardConsumer.task.run(); - } - } - } - synchronized (KinesisMessageDrivenChannelAdapter.this.consumerInvokers) { - // The attempt to survive if ShardConsumer has been added during synchronization - if (this.consumers.isEmpty()) { - KinesisMessageDrivenChannelAdapter.this.consumerInvokers.remove(this); - break; + else { + if (shardConsumer.task != null) { + shardConsumer.task.run(); } } } - } - finally { - this.lock.unlock(); + synchronized (KinesisMessageDrivenChannelAdapter.this.consumerInvokers) { + // The attempt to survive if ShardConsumer has been added during synchronization + if (this.consumers.isEmpty()) { + KinesisMessageDrivenChannelAdapter.this.consumerInvokers.remove(this); + break; + } + } } } 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 26ba70f..b67b739 100644 --- a/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java +++ b/src/test/java/org/springframework/integration/aws/inbound/KinesisMessageDrivenChannelAdapterTests.java @@ -150,8 +150,7 @@ public class KinesisMessageDrivenChannelAdapterTests { this.kinesisMessageDrivenChannelAdapter.stop(); this.kinesisMessageDrivenChannelAdapter.setListenerMode(ListenerMode.batch); - this.kinesisMessageDrivenChannelAdapter - .setCheckpointMode(CheckpointMode.record); + this.kinesisMessageDrivenChannelAdapter.setCheckpointMode(CheckpointMode.record); this.checkpointStore.put("SpringIntegration" + ":" + STREAM1 + ":" + "1", "1"); this.kinesisMessageDrivenChannelAdapter.start(); @@ -314,10 +313,12 @@ public class KinesisMessageDrivenChannelAdapterTests { adapter.setStartTimeout(10000); adapter.setDescribeStreamRetries(1); adapter.setConcurrency(10); + adapter.setRecordsLimit(25); DirectFieldAccessor dfa = new DirectFieldAccessor(adapter); dfa.setPropertyValue("describeStreamBackoff", 10); dfa.setPropertyValue("consumerBackoff", 10); + dfa.setPropertyValue("idleBetweenPolls", 1); return adapter; } @@ -373,10 +374,12 @@ public class KinesisMessageDrivenChannelAdapterTests { adapter.setOutputChannel(kinesisChannel()); adapter.setStartTimeout(10000); adapter.setDescribeStreamRetries(1); + adapter.setRecordsLimit(25); DirectFieldAccessor dfa = new DirectFieldAccessor(adapter); dfa.setPropertyValue("describeStreamBackoff", 10); dfa.setPropertyValue("consumerBackoff", 10); + dfa.setPropertyValue("idleBetweenPolls", 1); adapter.setConverter(String::new);