diff --git a/src/main/java/org/springframework/integration/aws/outbound/KinesisMessageHandler.java b/src/main/java/org/springframework/integration/aws/outbound/KinesisMessageHandler.java index 149799d..4b70b79 100644 --- a/src/main/java/org/springframework/integration/aws/outbound/KinesisMessageHandler.java +++ b/src/main/java/org/springframework/integration/aws/outbound/KinesisMessageHandler.java @@ -35,10 +35,13 @@ import org.springframework.messaging.Message; import org.springframework.util.Assert; import org.springframework.util.StringUtils; +import com.amazonaws.AmazonWebServiceRequest; import com.amazonaws.handlers.AsyncHandler; import com.amazonaws.services.kinesis.AmazonKinesisAsync; import com.amazonaws.services.kinesis.model.PutRecordRequest; import com.amazonaws.services.kinesis.model.PutRecordResult; +import com.amazonaws.services.kinesis.model.PutRecordsRequest; +import com.amazonaws.services.kinesis.model.PutRecordsResult; /** * The {@link AbstractMessageHandler} implementation for the Amazon Kinesis {@code putRecord(s)}. @@ -55,7 +58,7 @@ public class KinesisMessageHandler extends AbstractMessageHandler { private final AmazonKinesisAsync amazonKinesis; - private AsyncHandler asyncHandler; + private AsyncHandler asyncHandler; private Converter converter = new SerializingConverter(); @@ -77,7 +80,7 @@ public class KinesisMessageHandler extends AbstractMessageHandler { this.amazonKinesis = amazonKinesis; } - public void setAsyncHandler(AsyncHandler asyncHandler) { + public void setAsyncHandler(AsyncHandler asyncHandler) { this.asyncHandler = asyncHandler; } @@ -145,7 +148,40 @@ public class KinesisMessageHandler extends AbstractMessageHandler { } @Override + @SuppressWarnings("unchecked") protected void handleMessageInternal(Message message) throws Exception { + Future resultFuture = null; + if (message.getPayload() instanceof PutRecordsRequest) { + resultFuture = this.amazonKinesis.putRecordsAsync((PutRecordsRequest) message.getPayload(), + (AsyncHandler) this.asyncHandler); + } + else { + + PutRecordRequest putRecordRequest = (message.getPayload() instanceof PutRecordRequest) + ? (PutRecordRequest) message.getPayload() + : buildPutRecordRequest(message); + + resultFuture = this.amazonKinesis.putRecordAsync(putRecordRequest, + (AsyncHandler) this.asyncHandler); + } + + if (this.sync) { + Long sendTimeout = this.sendTimeoutExpression.getValue(this.evaluationContext, message, Long.class); + if (sendTimeout == null || sendTimeout < 0) { + resultFuture.get(); + } + else { + try { + resultFuture.get(sendTimeout, TimeUnit.MILLISECONDS); + } + catch (TimeoutException te) { + throw new MessageTimeoutException(message, "Timeout waiting for response from AmazonKinesis", te); + } + } + } + } + + private PutRecordRequest buildPutRecordRequest(Message message) { String stream = message.getHeaders().get(AwsHeaders.STREAM, String.class); if (!StringUtils.hasText(stream) && this.streamExpression != null) { stream = this.streamExpression.getValue(this.evaluationContext, message, String.class); @@ -172,29 +208,12 @@ public class KinesisMessageHandler extends AbstractMessageHandler { partitionKey = this.sequenceNumberExpression.getValue(this.evaluationContext, message, String.class); } - PutRecordRequest putRecordRequest = new PutRecordRequest() + return new PutRecordRequest() .withStreamName(stream) .withPartitionKey(partitionKey) .withExplicitHashKey(explicitHashKey) .withSequenceNumberForOrdering(sequenceNumber) .withData(ByteBuffer.wrap(this.converter.convert(message.getPayload()))); - - Future resultFuture = this.amazonKinesis.putRecordAsync(putRecordRequest, this.asyncHandler); - - if (this.sync) { - Long sendTimeout = this.sendTimeoutExpression.getValue(this.evaluationContext, message, Long.class); - if (sendTimeout == null || sendTimeout < 0) { - resultFuture.get(); - } - else { - try { - resultFuture.get(sendTimeout, TimeUnit.MILLISECONDS); - } - catch (TimeoutException te) { - throw new MessageTimeoutException(message, "Timeout waiting for response from AmazonKinesis", te); - } - } - } } } diff --git a/src/test/java/org/springframework/integration/aws/outbound/KinesisMessageHandlerTests.java b/src/test/java/org/springframework/integration/aws/outbound/KinesisMessageHandlerTests.java index f6905e4..e0cdefe 100644 --- a/src/test/java/org/springframework/integration/aws/outbound/KinesisMessageHandlerTests.java +++ b/src/test/java/org/springframework/integration/aws/outbound/KinesisMessageHandlerTests.java @@ -42,6 +42,7 @@ import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; import org.springframework.messaging.MessageHandler; import org.springframework.messaging.MessageHandlingException; +import org.springframework.messaging.support.GenericMessage; import org.springframework.messaging.support.MessageBuilder; import org.springframework.test.context.junit4.SpringRunner; @@ -49,6 +50,9 @@ import com.amazonaws.handlers.AsyncHandler; import com.amazonaws.services.kinesis.AmazonKinesisAsync; import com.amazonaws.services.kinesis.model.PutRecordRequest; import com.amazonaws.services.kinesis.model.PutRecordResult; +import com.amazonaws.services.kinesis.model.PutRecordsRequest; +import com.amazonaws.services.kinesis.model.PutRecordsRequestEntry; +import com.amazonaws.services.kinesis.model.PutRecordsResult; /** * @author Artem Bilan @@ -67,11 +71,12 @@ public class KinesisMessageHandlerTests { protected KinesisMessageHandler kinesisMessageHandler; @Autowired - protected AsyncHandler asyncHandler; + protected AsyncHandler asyncHandler; @Test + @SuppressWarnings("unchecked") public void testKinesisMessageHandler() { - Message message = MessageBuilder.withPayload("message").build(); + Message message = MessageBuilder.withPayload("message").build(); try { this.kinesisSendChannel.send(message); } @@ -100,7 +105,8 @@ public class KinesisMessageHandlerTests { ArgumentCaptor putRecordRequestArgumentCaptor = ArgumentCaptor.forClass(PutRecordRequest.class); - verify(this.amazonKinesis).putRecordAsync(putRecordRequestArgumentCaptor.capture(), eq(this.asyncHandler)); + verify(this.amazonKinesis).putRecordAsync(putRecordRequestArgumentCaptor.capture(), + eq((AsyncHandler) this.asyncHandler)); PutRecordRequest putRecordRequest = putRecordRequestArgumentCaptor.getValue(); @@ -109,6 +115,27 @@ public class KinesisMessageHandlerTests { assertThat(putRecordRequest.getSequenceNumberForOrdering()).isEqualTo("10"); assertThat(putRecordRequest.getExplicitHashKey()).isNull(); assertThat(putRecordRequest.getData()).isEqualTo(ByteBuffer.wrap("message".getBytes())); + + message = new GenericMessage<>(new PutRecordsRequest() + .withStreamName("myStream") + .withRecords(new PutRecordsRequestEntry() + .withData(ByteBuffer.wrap("test".getBytes())) + .withPartitionKey("testKey"))); + + this.kinesisSendChannel.send(message); + + ArgumentCaptor putRecordsRequestArgumentCaptor = + ArgumentCaptor.forClass(PutRecordsRequest.class); + verify(this.amazonKinesis).putRecordsAsync(putRecordsRequestArgumentCaptor.capture(), + eq((AsyncHandler) this.asyncHandler)); + + PutRecordsRequest putRecordsRequest = putRecordsRequestArgumentCaptor.getValue(); + + assertThat(putRecordsRequest.getStreamName()).isEqualTo("myStream"); + assertThat(putRecordsRequest.getRecords()) + .containsExactlyInAnyOrder(new PutRecordsRequestEntry() + .withData(ByteBuffer.wrap("test".getBytes())) + .withPartitionKey("testKey")); } @@ -120,14 +147,19 @@ public class KinesisMessageHandlerTests { @SuppressWarnings("unchecked") public AmazonKinesisAsync amazonKinesis() { AmazonKinesisAsync mock = mock(AmazonKinesisAsync.class); + given(mock.putRecordAsync(any(PutRecordRequest.class), any(AsyncHandler.class))) .willReturn(mock(Future.class)); + + given(mock.putRecordsAsync(any(PutRecordsRequest.class), any(AsyncHandler.class))) + .willReturn(mock(Future.class)); + return mock; } @Bean @SuppressWarnings("unchecked") - public AsyncHandler asyncHandler() { + public AsyncHandler asyncHandler() { return mock(AsyncHandler.class); }