GH-51: Add headers mapping to channel adapters

Fixes spring-projects/spring-integration-aws#51
This commit is contained in:
Artem Bilan
2018-04-06 16:51:47 -04:00
parent 2b57f34801
commit 6a999b2873
14 changed files with 464 additions and 29 deletions

View File

@@ -17,6 +17,7 @@
package org.springframework.integration.aws.kinesis;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.entry;
import java.util.Date;
import java.util.HashSet;
@@ -42,6 +43,7 @@ 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.integration.support.json.EmbeddedJsonHeadersMessageMapper;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.MessageHandler;
@@ -95,11 +97,13 @@ public class KinesisIntegrationTests {
this.kinesisSendChannel.send(
MessageBuilder.withPayload(now)
.setHeader(AwsHeaders.STREAM, TEST_STREAM)
.setHeader("foo", "BAR")
.build());
Message<?> receive = this.kinesisReceiveChannel.receive(10_000);
assertThat(receive).isNotNull();
assertThat(receive.getPayload()).isEqualTo(now);
assertThat(receive.getHeaders()).contains(entry("foo", "BAR"));
Message<?> errorMessage = this.errorChannel.receive(10_000);
assertThat(errorMessage).isNotNull();
@@ -141,6 +145,7 @@ public class KinesisIntegrationTests {
public MessageHandler kinesisMessageHandler() {
KinesisMessageHandler kinesisMessageHandler = new KinesisMessageHandler(KINESIS_LOCAL_RUNNING.getKinesis());
kinesisMessageHandler.setPartitionKey("1");
kinesisMessageHandler.setEmbeddedHeadersMapper(new EmbeddedJsonHeadersMessageMapper("foo"));
return kinesisMessageHandler;
}
@@ -156,6 +161,7 @@ public class KinesisIntegrationTests {
adapter.setErrorChannel(errorChannel());
adapter.setErrorMessageStrategy(new KinesisMessageHeaderErrorMessageStrategy());
adapter.setCheckpointStore(checkpointStore());
adapter.setEmbeddedHeadersMapper(new EmbeddedJsonHeadersMessageMapper("foo"));
return adapter;
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2016-2017 the original author or authors.
* Copyright 2016-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.
@@ -17,6 +17,7 @@
package org.springframework.integration.aws.outbound;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.entry;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.BDDMockito.given;
@@ -38,6 +39,7 @@ import org.springframework.core.serializer.support.SerializingConverter;
import org.springframework.integration.annotation.ServiceActivator;
import org.springframework.integration.aws.support.AwsHeaders;
import org.springframework.integration.config.EnableIntegration;
import org.springframework.integration.support.json.EmbeddedJsonHeadersMessageMapper;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.MessageHandler;
@@ -77,7 +79,7 @@ public class KinesisMessageHandlerTests {
@Test
@SuppressWarnings("unchecked")
public void testKinesisMessageHandler() {
public void testKinesisMessageHandler() throws Exception {
Message<?> message = MessageBuilder.withPayload("message").build();
try {
this.kinesisSendChannel.send(message);
@@ -101,6 +103,7 @@ public class KinesisMessageHandlerTests {
message = MessageBuilder.fromMessage(message)
.setHeader(AwsHeaders.PARTITION_KEY, "fooKey")
.setHeader(AwsHeaders.SEQUENCE_NUMBER, "10")
.setHeader("foo", "bar")
.build();
this.kinesisSendChannel.send(message);
@@ -119,7 +122,12 @@ public class KinesisMessageHandlerTests {
assertThat(putRecordRequest.getPartitionKey()).isEqualTo("fooKey");
assertThat(putRecordRequest.getSequenceNumberForOrdering()).isEqualTo("10");
assertThat(putRecordRequest.getExplicitHashKey()).isNull();
assertThat(putRecordRequest.getData()).isEqualTo(ByteBuffer.wrap("message".getBytes()));
Message<?> messageToCheck = new EmbeddedJsonHeadersMessageMapper()
.toMessage(putRecordRequest.getData().array());
assertThat(messageToCheck.getHeaders()).contains(entry("foo", "bar"));
assertThat(messageToCheck.getPayload()).isEqualTo("message".getBytes());
AsyncHandler<?, ?> asyncHandler = asyncHandlerArgumentCaptor.getValue();
@@ -195,6 +203,7 @@ public class KinesisMessageHandlerTests {
}
});
kinesisMessageHandler.setEmbeddedHeadersMapper(new EmbeddedJsonHeadersMessageMapper("foo"));
return kinesisMessageHandler;
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2016-2017 the original author or authors.
* Copyright 2016-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.
@@ -22,6 +22,8 @@ import static org.mockito.BDDMockito.willAnswer;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import java.util.Map;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.mockito.ArgumentCaptor;
@@ -33,12 +35,14 @@ import org.springframework.expression.spel.standard.SpelExpressionParser;
import org.springframework.integration.annotation.ServiceActivator;
import org.springframework.integration.aws.support.AwsHeaders;
import org.springframework.integration.aws.support.SnsBodyBuilder;
import org.springframework.integration.aws.support.SnsHeaderMapper;
import org.springframework.integration.channel.QueueChannel;
import org.springframework.integration.config.EnableIntegration;
import org.springframework.integration.support.MessageBuilder;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.MessageHandler;
import org.springframework.messaging.MessageHeaders;
import org.springframework.messaging.PollableChannel;
import org.springframework.scheduling.annotation.AsyncResult;
import org.springframework.test.annotation.DirtiesContext;
@@ -47,6 +51,7 @@ import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
import com.amazonaws.handlers.AsyncHandler;
import com.amazonaws.services.sns.AmazonSNSAsync;
import com.amazonaws.services.sns.model.MessageAttributeValue;
import com.amazonaws.services.sns.model.PublishRequest;
import com.amazonaws.services.sns.model.PublishResult;
@@ -78,6 +83,7 @@ public class SnsMessageHandlerTests {
Message<?> message = MessageBuilder.withPayload(payload)
.setHeader("topic", "topic")
.setHeader("subject", "subject")
.setHeader("foo", "bar")
.build();
this.sendToSnsChannel.send(message);
@@ -96,6 +102,13 @@ public class SnsMessageHandlerTests {
assertThat(publishRequest.getMessage())
.isEqualTo("{\"default\":\"foo\",\"sms\":\"{\\\"foo\\\" : \\\"bar\\\"}\"}");
Map<String, MessageAttributeValue> messageAttributes = publishRequest.getMessageAttributes();
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.ID);
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.TIMESTAMP);
assertThat(messageAttributes).containsKey("foo");
assertThat(messageAttributes.get("foo").getStringValue()).isEqualTo("bar");
assertThat(reply.getHeaders().get(AwsHeaders.MESSAGE_ID)).isEqualTo("111");
assertThat(reply.getHeaders().get(AwsHeaders.TOPIC)).isEqualTo("topic");
assertThat(reply.getPayload()).isSameAs(payload);
@@ -135,6 +148,9 @@ public class SnsMessageHandlerTests {
snsMessageHandler.setSubjectExpression(PARSER.parseExpression("headers.subject"));
snsMessageHandler.setBodyExpression(PARSER.parseExpression("payload"));
snsMessageHandler.setOutputChannel(resultChannel());
SnsHeaderMapper headerMapper = new SnsHeaderMapper();
headerMapper.setOutboundHeaderNames("foo");
snsMessageHandler.setHeaderMapper(headerMapper);
return snsMessageHandler;
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2015-2017 the original author or authors.
* Copyright 2015-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.
@@ -23,6 +23,8 @@ import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import java.util.Map;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.mockito.ArgumentCaptor;
@@ -39,6 +41,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.MessageHeaders;
import org.springframework.messaging.support.MessageBuilder;
import org.springframework.test.context.ContextConfiguration;
import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
@@ -47,6 +50,7 @@ import com.amazonaws.handlers.AsyncHandler;
import com.amazonaws.services.sqs.AmazonSQSAsync;
import com.amazonaws.services.sqs.model.GetQueueUrlRequest;
import com.amazonaws.services.sqs.model.GetQueueUrlResult;
import com.amazonaws.services.sqs.model.MessageAttributeValue;
import com.amazonaws.services.sqs.model.SendMessageRequest;
/**
@@ -105,8 +109,18 @@ public class SqsMessageHandlerTests {
verify(this.amazonSqs, times(3))
.sendMessageAsync(sendMessageRequestArgumentCaptor.capture(), any(AsyncHandler.class));
assertThat(sendMessageRequestArgumentCaptor.getValue().getQueueUrl())
SendMessageRequest sendMessageRequestArgumentCaptorValue = sendMessageRequestArgumentCaptor.getValue();
assertThat(sendMessageRequestArgumentCaptorValue.getQueueUrl())
.isEqualTo("http://queue-url.com/baz");
Map<String, MessageAttributeValue> messageAttributes =
sendMessageRequestArgumentCaptorValue.getMessageAttributes();
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.ID);
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.TIMESTAMP);
assertThat(messageAttributes).containsKey("foo");
assertThat(messageAttributes.get("foo").getStringValue()).isEqualTo("baz");
}
@Configuration