Test SqsMessageHandler against Localstack

* Fix `SqsMessageHandler` to check for `sequenceNumber` before populating it
to replay headers
This commit is contained in:
abilan
2023-03-15 12:32:44 -04:00
parent f13ef11798
commit d003e1b850
3 changed files with 79 additions and 80 deletions

View File

@@ -236,8 +236,13 @@ public class SqsMessageHandler extends AbstractAwsMessageHandler<Map<String, Mes
@Override
protected Map<String, ?> additionalOnSuccessHeaders(AwsRequest request, AwsResponse response) {
if (response instanceof SendMessageResponse sendMessageResponse) {
return Map.of(AwsHeaders.MESSAGE_ID, sendMessageResponse.messageId(),
AwsHeaders.SEQUENCE_NUMBER, sendMessageResponse.sequenceNumber());
Map<String, Object> headers = new HashMap<>();
headers.put(AwsHeaders.MESSAGE_ID, sendMessageResponse.messageId());
String sequenceNumber = sendMessageResponse.sequenceNumber();
if (StringUtils.hasText(sequenceNumber)) {
headers.put(AwsHeaders.SEQUENCE_NUMBER, sequenceNumber);
}
return headers;
}
return null;
}

View File

@@ -29,6 +29,7 @@ import software.amazon.awssdk.services.cloudwatch.CloudWatchAsyncClient;
import software.amazon.awssdk.services.dynamodb.DynamoDbAsyncClient;
import software.amazon.awssdk.services.kinesis.KinesisAsyncClient;
import software.amazon.awssdk.services.s3.S3AsyncClient;
import software.amazon.awssdk.services.sqs.SqsAsyncClient;
/**
* The base contract for JUnit tests based on the container for Localstack.
@@ -51,7 +52,8 @@ public interface LocalstackContainerTest {
LocalStackContainer.Service.DYNAMODB,
LocalStackContainer.Service.KINESIS,
LocalStackContainer.Service.CLOUDWATCH,
LocalStackContainer.Service.S3);
LocalStackContainer.Service.S3,
LocalStackContainer.Service.SQS);
@BeforeAll
static void startContainer() {
@@ -74,6 +76,10 @@ public interface LocalstackContainerTest {
return applyAwsClientOptions(S3AsyncClient.builder(), LocalStackContainer.Service.CLOUDWATCH);
}
static SqsAsyncClient sqsClient() {
return applyAwsClientOptions(SqsAsyncClient.builder(), LocalStackContainer.Service.SQS);
}
static AwsCredentialsProvider credentialsProvider() {
return StaticCredentialsProvider.create(
AwsBasicCredentials.create(LOCAL_STACK_CONTAINER.getAccessKey(), LOCAL_STACK_CONTAINER.getSecretKey()));
@@ -88,5 +94,4 @@ public interface LocalstackContainerTest {
.endpointOverride(LOCAL_STACK_CONTAINER.getEndpointOverride(serviceToBuild))
.build();
}
}

View File

@@ -18,20 +18,15 @@ package org.springframework.integration.aws.outbound;
import java.util.Map;
import java.util.concurrent.CompletableFuture;
import java.util.function.Consumer;
import java.util.concurrent.atomic.AtomicReference;
import io.awspring.cloud.sqs.listener.QueueNotFoundStrategy;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import software.amazon.awssdk.services.sqs.SqsAsyncClient;
import software.amazon.awssdk.services.sqs.model.CreateQueueRequest;
import software.amazon.awssdk.services.sqs.model.CreateQueueResponse;
import software.amazon.awssdk.services.sqs.model.GetQueueUrlRequest;
import software.amazon.awssdk.services.sqs.model.GetQueueUrlResponse;
import software.amazon.awssdk.services.sqs.model.MessageAttributeValue;
import software.amazon.awssdk.services.sqs.model.SendMessageRequest;
import software.amazon.awssdk.services.sqs.model.SendMessageResponse;
import software.amazon.awssdk.services.sqs.model.QueueAttributeName;
import software.amazon.awssdk.services.sqs.model.ReceiveMessageResponse;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.annotation.Bean;
@@ -39,6 +34,7 @@ import org.springframework.context.annotation.Configuration;
import org.springframework.expression.Expression;
import org.springframework.expression.spel.standard.SpelExpressionParser;
import org.springframework.integration.annotation.ServiceActivator;
import org.springframework.integration.aws.LocalstackContainerTest;
import org.springframework.integration.aws.support.AwsHeaders;
import org.springframework.integration.config.EnableIntegration;
import org.springframework.messaging.Message;
@@ -47,17 +43,10 @@ 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.annotation.DirtiesContext;
import org.springframework.test.context.junit.jupiter.SpringJUnitConfig;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatExceptionOfType;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.BDDMockito.given;
import static org.mockito.BDDMockito.willAnswer;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
/**
* Instantiating SqsMessageHandler using amazonSqs.
@@ -66,13 +55,16 @@ import static org.mockito.Mockito.verify;
* @author Rahul Pilani
* @author Seth Kelly
*/
@Disabled("Revise in favor of Local Stack")
@SpringJUnitConfig
@DirtiesContext(classMode = DirtiesContext.ClassMode.AFTER_EACH_TEST_METHOD)
public class SqsMessageHandlerTests {
public class SqsMessageHandlerTests implements LocalstackContainerTest {
@Autowired
protected SqsAsyncClient amazonSqs;
private static final AtomicReference<String> fooUrl = new AtomicReference<>();
private static final AtomicReference<String> barUrl = new AtomicReference<>();
private static final AtomicReference<String> bazUrl = new AtomicReference<>();
private static SqsAsyncClient AMAZON_SQS;
@Autowired
protected MessageChannel sqsSendChannel;
@@ -83,8 +75,21 @@ public class SqsMessageHandlerTests {
@Autowired
protected SqsMessageHandler sqsMessageHandler;
@Autowired
protected SqsMessageHandler sqsMessageHandlerWithQueueAutoCreate;
@BeforeAll
static void setup() {
AMAZON_SQS = LocalstackContainerTest.sqsClient();
CompletableFuture<?> foo =
AMAZON_SQS.createQueue(request -> request.queueName("foo"))
.thenAccept(response -> fooUrl.set(response.queueUrl()));
CompletableFuture<?> bar =
AMAZON_SQS.createQueue(request -> request.queueName("bar"))
.thenAccept(response -> barUrl.set(response.queueUrl()));
CompletableFuture<?> baz =
AMAZON_SQS.createQueue(request -> request.queueName("baz"))
.thenAccept(response -> bazUrl.set(response.queueUrl()));
CompletableFuture.allOf(foo, bar, baz).join();
}
@Test
void testSqsMessageHandler() {
@@ -96,28 +101,43 @@ public class SqsMessageHandlerTests {
this.sqsMessageHandler.setQueue("foo");
this.sqsSendChannel.send(message);
ArgumentCaptor<SendMessageRequest> sendMessageRequestArgumentCaptor =
ArgumentCaptor.forClass(SendMessageRequest.class);
verify(this.amazonSqs).sendMessage(sendMessageRequestArgumentCaptor.capture());
assertThat(sendMessageRequestArgumentCaptor.getValue().queueUrl()).isEqualTo("https://queue-url.com/foo");
ReceiveMessageResponse receiveMessageResponse =
AMAZON_SQS.receiveMessage(request -> request.queueUrl(fooUrl.get()).waitTimeSeconds(10))
.join();
assertThat(receiveMessageResponse.hasMessages()).isTrue();
assertThat(receiveMessageResponse.messages().get(0).body()).isEqualTo("message");
Message<String> message2 = MessageBuilder.withPayload("message").setHeader(AwsHeaders.QUEUE, "bar").build();
this.sqsSendChannel.send(message2);
verify(this.amazonSqs, times(2)).sendMessage(sendMessageRequestArgumentCaptor.capture());
assertThat(sendMessageRequestArgumentCaptor.getValue().queueUrl()).isEqualTo("https://queue-url.com/bar");
receiveMessageResponse =
AMAZON_SQS.receiveMessage(request -> request.queueUrl(barUrl.get()).waitTimeSeconds(10))
.join();
assertThat(receiveMessageResponse.hasMessages()).isTrue();
assertThat(receiveMessageResponse.messages().get(0).body()).isEqualTo("message");
SpelExpressionParser spelExpressionParser = new SpelExpressionParser();
Expression expression = spelExpressionParser.parseExpression("headers.foo");
this.sqsMessageHandler.setQueueExpression(expression);
message2 = MessageBuilder.withPayload("message").setHeader("foo", "baz").build();
this.sqsSendChannel.send(message2);
verify(this.amazonSqs, times(3)).sendMessage(sendMessageRequestArgumentCaptor.capture());
SendMessageRequest sendMessageRequestArgumentCaptorValue = sendMessageRequestArgumentCaptor.getValue();
assertThat(sendMessageRequestArgumentCaptorValue.queueUrl()).isEqualTo("https://queue-url.com/baz");
receiveMessageResponse =
AMAZON_SQS.receiveMessage(request ->
request.queueUrl(bazUrl.get())
.messageAttributeNames(QueueAttributeName.ALL.toString())
.waitTimeSeconds(10))
.join();
Map<String, MessageAttributeValue> messageAttributes = sendMessageRequestArgumentCaptorValue.messageAttributes();
assertThat(receiveMessageResponse.hasMessages()).isTrue();
software.amazon.awssdk.services.sqs.model.Message message1 = receiveMessageResponse.messages().get(0);
assertThat(message1.body()).isEqualTo("message");
Map<String, MessageAttributeValue> messageAttributes = message1.messageAttributes();
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.ID);
assertThat(messageAttributes).doesNotContainKey(MessageHeaders.TIMESTAMP);
@@ -126,69 +146,38 @@ public class SqsMessageHandlerTests {
}
@Test
@SuppressWarnings("unchecked")
void testSqsMessageHandlerWithAutoQueueCreate() {
Message<String> message = MessageBuilder.withPayload("message").build();
this.sqsMessageHandlerWithQueueAutoCreate.setQueue("foo");
this.sqsSendChannelWithAutoCreate.send(message);
ArgumentCaptor<CreateQueueRequest> createQueueRequestArgumentCaptor =
ArgumentCaptor.forClass(CreateQueueRequest.class);
verify(this.amazonSqs).createQueue(createQueueRequestArgumentCaptor.capture());
assertThat(createQueueRequestArgumentCaptor.getValue().queueName()).isEqualTo("foo");
ArgumentCaptor<SendMessageRequest> sendMessageRequestArgumentCaptor =
ArgumentCaptor.forClass(SendMessageRequest.class);
verify(this.amazonSqs).sendMessage(sendMessageRequestArgumentCaptor.capture());
assertThat(sendMessageRequestArgumentCaptor.getValue().queueUrl()).isEqualTo("https://queue-url.com/foo");
ReceiveMessageResponse autoCreateQueueResponse =
AMAZON_SQS.getQueueUrl(request -> request.queueName("autoCreateQueue"))
.thenCompose(response ->
AMAZON_SQS.receiveMessage(request ->
request.queueUrl(response.queueUrl()).waitTimeSeconds(10)))
.join();
assertThat(autoCreateQueueResponse.hasMessages()).isTrue();
assertThat(autoCreateQueueResponse.messages().get(0).body()).isEqualTo("message");
}
@Configuration
@EnableIntegration
public static class ContextConfiguration {
@Bean
@SuppressWarnings("unchecked")
public SqsAsyncClient amazonSqs() {
SqsAsyncClient amazonSqs = mock(SqsAsyncClient.class);
willAnswer(invocation -> {
GetQueueUrlRequest getQueueUrlRequest = (GetQueueUrlRequest) invocation.getArguments()[0];
return CompletableFuture.completedFuture(
GetQueueUrlResponse.builder()
.queueUrl("https://queue-url.com/" + getQueueUrlRequest.queueName())
.build());
}).given(amazonSqs).getQueueUrl(any(GetQueueUrlRequest.class));
willAnswer(invocation -> {
CreateQueueRequest createQueueRequest = (CreateQueueRequest) invocation.getArguments()[0];
return CompletableFuture.completedFuture(
CreateQueueResponse.builder()
.queueUrl("https://queue-url.com/" + createQueueRequest.queueName())
.build());
}).given(amazonSqs).createQueue(any(Consumer.class));
given(amazonSqs.sendMessage(any(SendMessageRequest.class)))
.willReturn(CompletableFuture.completedFuture(
SendMessageResponse.builder()
.messageId("testId")
.sequenceNumber("1")
.build()));
return amazonSqs;
}
@Bean
@ServiceActivator(inputChannel = "sqsSendChannel")
public MessageHandler sqsMessageHandler() {
return new SqsMessageHandler(amazonSqs());
return new SqsMessageHandler(AMAZON_SQS);
}
@Bean
@ServiceActivator(inputChannel = "sqsSendChannelWithAutoCreate")
public MessageHandler sqsMessageHandlerWithQueueAutoCreate() {
SqsMessageHandler sqsMessageHandler = new SqsMessageHandler(amazonSqs());
SqsMessageHandler sqsMessageHandler = new SqsMessageHandler(AMAZON_SQS);
sqsMessageHandler.setQueueNotFoundStrategy(QueueNotFoundStrategy.CREATE);
sqsMessageHandler.setQueue("autoCreateQueue");
return sqsMessageHandler;
}