diff --git a/docs/src/main/asciidoc/sqs.adoc b/docs/src/main/asciidoc/sqs.adoc index d9d91c011..58ed1234d 100644 --- a/docs/src/main/asciidoc/sqs.adoc +++ b/docs/src/main/asciidoc/sqs.adoc @@ -319,8 +319,10 @@ SendResult.Batch sendMany(String queue, Collection> messages); ``` NOTE: To send a collection of objects, it is recommended to use `sendMany(String queue, Collection> messages)` to optimize throughput. -To send a collection of objects in a single message, the collection must be wrapped in an object. The underlying AWS SQS API has a limitation that only up to 10 messages can be sent in a single batch request using `SendMessageBatch`. If more than 10 messages are passed to `sendMany()`, the AWS SDK will throw a `TooManyEntriesInBatchRequestException`. -This limitation is documented in the https://docs.aws.amazon.com/AWSSimpleQueueService/latest/APIReference/API_SendMessageBatch.html[AWS SQS API Reference for SendMessageBatch]. As of now, Spring Cloud AWS does not automatically split larger collections into smaller batches of 10 or fewer messages. Users are responsible for ensuring the batch size complies with this AWS limit. +To send a collection of objects in a single message, the collection must be wrapped in an object. +The underlying AWS SQS API has a limitation that only up to 10 messages can be sent in a single batch request using `SendMessageBatch`. +Since 4.2.0, Spring Cloud AWS automatically splits larger collections into batches of 10 or fewer messages. +For standard queues, all batches are sent in parallel. For FIFO queues, batches are sent sequentially to preserve the message order, with messages grouped by message group ID. An example using the `options` variant follows: diff --git a/spring-cloud-aws-samples/spring-cloud-aws-sqs-sample/src/main/java/io/awspring/cloud/sqs/sample/SendManyBatchSample.java b/spring-cloud-aws-samples/spring-cloud-aws-sqs-sample/src/main/java/io/awspring/cloud/sqs/sample/SendManyBatchSample.java new file mode 100644 index 000000000..03b8582ef --- /dev/null +++ b/spring-cloud-aws-samples/spring-cloud-aws-sqs-sample/src/main/java/io/awspring/cloud/sqs/sample/SendManyBatchSample.java @@ -0,0 +1,61 @@ +/* + * Copyright 2013-2025 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package io.awspring.cloud.sqs.sample; + +import io.awspring.cloud.sqs.annotation.SqsListener; +import io.awspring.cloud.sqs.operations.SendResult; +import io.awspring.cloud.sqs.operations.SqsTemplate; +import java.util.List; +import java.util.stream.IntStream; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import org.springframework.boot.ApplicationRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.messaging.Message; +import org.springframework.messaging.support.MessageBuilder; + +/** + * Sample demonstrating {@link SqsTemplate#sendMany} sending more than 10 messages at once. The template automatically + * partitions the messages into batches of 10 and sends them in parallel (for standard queues) or sequentially per + * message group (for FIFO queues). + * + * @author José Iêdo + */ +@Configuration +public class SendManyBatchSample { + + private static final Logger LOGGER = LoggerFactory.getLogger(SendManyBatchSample.class); + + private static final String QUEUE_NAME = "send-many-batch-queue"; + + @SqsListener(queueNames = QUEUE_NAME, maxMessagesPerPoll = "25", maxConcurrentMessages = "25") + void listen(List> messages) { + LOGGER.info("Received {} messages: {}", messages.size(), messages.stream().map(Message::getPayload).toList()); + } + + @Bean + public ApplicationRunner sendManyMessages(SqsTemplate sqsTemplate) { + return args -> { + List> messages = IntStream.range(0, 25).mapToObj(index -> "Message-" + index) + .map(payload -> MessageBuilder.withPayload(payload).build()).toList(); + LOGGER.info("Sending {} messages to queue {}", messages.size(), QUEUE_NAME); + SendResult.Batch result = sqsTemplate.sendMany(QUEUE_NAME, messages); + LOGGER.info("Sent successfully: {}, failed: {}", result.successful().size(), result.failed().size()); + }; + } + +} diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplate.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplate.java index d889669cd..b6b21fe1b 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplate.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplate.java @@ -16,6 +16,7 @@ package io.awspring.cloud.sqs.operations; import io.awspring.cloud.core.support.JacksonPresent; +import io.awspring.cloud.sqs.CollectionUtils; import io.awspring.cloud.sqs.FifoUtils; import io.awspring.cloud.sqs.MessageHeaderUtils; import io.awspring.cloud.sqs.QueueAttributesResolver; @@ -35,9 +36,11 @@ import io.awspring.cloud.sqs.support.converter.legacy.LegacyJackson2SqsMessagingMessageConverter; import io.awspring.cloud.sqs.support.observation.SqsTemplateObservation; import java.time.Duration; +import java.util.ArrayList; import java.util.Collection; import java.util.Collections; import java.util.HashMap; +import java.util.List; import java.util.Map; import java.util.Optional; import java.util.UUID; @@ -80,11 +83,14 @@ * @author Zhong Xi Lu * @author Hyunggeol Lee * @author Jeongmin Kim + * @author José Iêdo * * @since 3.0 */ public class SqsTemplate extends AbstractMessagingTemplate implements SqsOperations, SqsAsyncOperations { + private static final int SQS_MAX_BATCH_SIZE = 10; + private static final Logger logger = LoggerFactory.getLogger(SqsTemplate.class); private static final SqsTemplateObservation.SqsSpecifics SQS_OBSERVATION_SPECIFICS = new SqsTemplateObservation.SqsSpecifics(); @@ -365,13 +371,158 @@ private SendMessageRequest doCreateSendMessageRequest(Message message, QueueAttr .messageSystemAttributes(mapMessageSystemAttributes(message)).build(); } + /** + * Sends a collection of messages using one or more SQS batch requests. + *

+ * The provided messages are automatically partitioned into batches of up to 10 messages, which is the maximum size + * supported by Amazon SQS. + *

+ * For standard queues, all batches are sent in parallel. + *

+ * For FIFO queues, messages are first grouped by + * {@link io.awspring.cloud.sqs.listener.SqsHeaders.MessageSystemAttributes#SQS_MESSAGE_GROUP_ID_HEADER message + * group ID}. Groups larger than 10 messages are sent sequentially to preserve message ordering within each group, + * with a skip-on-failure strategy: if a batch completes with a partial failure, no subsequent batches for that + * group are sent. + *

+ * Groups with up to 10 messages are bin-packed into shared batches on a best-effort basis (first-fit decreasing), + * reducing the number of requests while keeping each group whole within a single batch to preserve ordering. Packed + * batches are sent in parallel, as are large-group chains across different groups. + */ @Override protected CompletableFuture> doSendBatchAsync(String endpointName, Collection messages, Collection> originalMessages) { logger.debug("Sending messages {} to endpoint {}", messages, endpointName); + Map> originalMessagesById = originalMessages.stream() + .collect(Collectors.toMap(MessageHeaderUtils::getRawMessageId, msg -> msg)); + if (messages.size() <= SQS_MAX_BATCH_SIZE) { + return sendSingleBatch(endpointName, messages, originalMessagesById); + } + return FifoUtils.isFifo(endpointName) ? sendFifoBatches(endpointName, messages, originalMessagesById) + : sendStandardBatches(endpointName, messages, originalMessagesById); + } + + private CompletableFuture> sendSingleBatch(String endpointName, + Collection messages, Map> originalMessagesById) { return createSendMessageBatchRequest(endpointName, messages).thenCompose(this.sqsAsyncClient::sendMessageBatch) - .thenApply(response -> createSendResultBatch(response, endpointName, originalMessages.stream() - .collect(Collectors.toMap(MessageHeaderUtils::getRawMessageId, msg -> msg)))); + .thenApply(response -> createSendResultBatch(response, endpointName, originalMessagesById)); + } + + private CompletableFuture> sendPartitionedBatch(String endpointName, + Collection messages, Map> originalMessagesById) { + return sendSingleBatch(endpointName, messages, originalMessagesById) + .exceptionally(t -> createFailedBatchResult(messages, t, endpointName, originalMessagesById)); + } + + private SendResult.Batch createFailedBatchResult(Collection partition, Throwable throwable, + String endpointName, Map> originalMessagesById) { + Throwable cause = throwable; + if (cause instanceof java.util.concurrent.CompletionException) { + cause = cause.getCause(); + } + String errorMessage = cause != null && cause.getMessage() != null ? cause.getMessage() : "Unknown error"; + Map additionalInformation = Map.of(SqsTemplateParameters.EXCEPTION_PARAMETER_NAME, cause); + List> failed = partition.stream().map(msg -> new SendResult.Failed<>(errorMessage, + endpointName, originalMessagesById.get(msg.messageId()), additionalInformation)).toList(); + return new SendResult.Batch<>(List.of(), failed); + } + + private CompletableFuture> sendStandardBatches(String endpointName, + Collection messages, Map> originalMessagesById) { + List>> futures = CollectionUtils.partition(messages, SQS_MAX_BATCH_SIZE) + .stream().map(partition -> sendPartitionedBatch(endpointName, partition, originalMessagesById)) + .toList(); + return combineBatchFutures(futures); + } + + private CompletableFuture> sendFifoBatches(String endpointName, + Collection messages, Map> originalMessagesById) { + Map> groupedByMessageGroup = messages.stream().collect(Collectors.groupingBy(msg -> { + String groupId = msg.attributes().get(MessageSystemAttributeName.MESSAGE_GROUP_ID); + return groupId != null ? groupId : ""; + })); + Map>> partitioned = groupedByMessageGroup.values().stream() + .collect(Collectors.partitioningBy(group -> group.size() <= SQS_MAX_BATCH_SIZE)); + List> smallGroups = partitioned.get(true); + List> largeGroups = partitioned.get(false); + List>> futures = largeGroups.stream() + .map(msgs -> sendSequentialBatches(endpointName, msgs, originalMessagesById)) + .collect(Collectors.toList()); + if (!smallGroups.isEmpty()) { + binPackSmallFifoGroups(smallGroups, SQS_MAX_BATCH_SIZE).stream() + .map(batch -> sendPartitionedBatch(endpointName, batch, originalMessagesById)) + .forEach(futures::add); + } + return combineBatchFutures(futures); + } + + /** + * Bin-pack small FIFO groups into shared batches using first-fit decreasing algorithm. Each group is kept whole + * within a single batch. Groups are sorted by size descending before packing to minimize the number of batches. + * @param smallGroups groups with size <= maxBatchSize + * @param maxBatchSize the maximum number of messages per batch (SQS limit is 10) + * @return packed batches, each containing one or more whole groups + */ + private static List> binPackSmallFifoGroups(List> smallGroups, int maxBatchSize) { + Assert.notNull(smallGroups, "smallGroups must not be null"); + Assert.isTrue(maxBatchSize > 0, "maxBatchSize must be positive"); + smallGroups.sort((a, b) -> Integer.compare(b.size(), a.size())); + List> packedBatches = new ArrayList<>(); + for (List group : smallGroups) { + boolean packed = false; + for (List batch : packedBatches) { + if (batch.size() + group.size() <= maxBatchSize) { + batch.addAll(group); + packed = true; + break; + } + } + if (!packed) { + packedBatches.add(new ArrayList<>(group)); + } + } + return packedBatches; + } + + private CompletableFuture> combineBatchFutures( + List>> futures) { + return CompletableFuture.allOf(futures.toArray(CompletableFuture[]::new)) + .thenApply(v -> futures.stream().map(CompletableFuture::join) + .reduce(new SendResult.Batch<>(List.of(), List.of()), this::mergeBatchResults)); + } + + private CompletableFuture> sendSequentialBatches(String endpointName, + List messages, Map> originalMessagesById) { + CompletableFuture> result = CompletableFuture + .completedFuture(new SendResult.Batch<>(List.of(), List.of())); + for (Collection partition : CollectionUtils.partition(messages, SQS_MAX_BATCH_SIZE)) { + result = result.thenCompose(acc -> { + if (!acc.failed().isEmpty()) { + return CompletableFuture.completedFuture( + mergeBatchResults(acc, createSkippedResult(partition, endpointName, originalMessagesById))); + } + return sendPartitionedBatch(endpointName, partition, originalMessagesById) + .thenApply(batchResult -> mergeBatchResults(acc, batchResult)); + }); + } + return result; + } + + private SendResult.Batch createSkippedResult(Collection partition, String endpointName, + Map> originalMessagesById) { + List> skipped = partition.stream() + .map(msg -> new SendResult.Failed<>("Skipped due to previous batch failure", endpointName, + originalMessagesById.get(msg.messageId()), Map.of())) + .toList(); + return new SendResult.Batch<>(List.of(), skipped); + } + + private SendResult.Batch mergeBatchResults(SendResult.Batch batch1, SendResult.Batch batch2) { + List> allSuccessful = new ArrayList<>(batch1.successful()); + allSuccessful.addAll(batch2.successful()); + List> allFailed = new ArrayList<>(batch1.failed()); + allFailed.addAll(batch2.failed()); + return new SendResult.Batch<>(allSuccessful, allFailed); } private SendResult.Batch createSendResultBatch(SendMessageBatchResponse response, String endpointName, @@ -937,8 +1088,8 @@ public SqsReceiveOptionsImpl visibilityTimeout(Duration visibilityTimeout) { @Override public SqsReceiveOptionsImpl maxNumberOfMessages(Integer maxNumberOfMessages) { Assert.notNull(maxNumberOfMessages, "maxNumberOfMessages must not be null"); - Assert.isTrue(maxNumberOfMessages > 0 && maxNumberOfMessages <= 10, - "maxNumberOfMessages must be between 0 and 10"); + Assert.isTrue(maxNumberOfMessages > 0 && maxNumberOfMessages <= SQS_MAX_BATCH_SIZE, + "maxNumberOfMessages must be between 0 and " + SQS_MAX_BATCH_SIZE); this.maxNumberOfMessages = maxNumberOfMessages; return this; } diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplateParameters.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplateParameters.java index 8ae127ed4..2a3179224 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplateParameters.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/operations/SqsTemplateParameters.java @@ -44,4 +44,9 @@ public class SqsTemplateParameters { */ public static final String ERROR_CODE_PARAMETER_NAME = "code"; + /** + * The exception that was thrown. + */ + public static final String EXCEPTION_PARAMETER_NAME = "exception"; + } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsFifoIntegrationTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsFifoIntegrationTests.java index 71f7c67bf..62d3f7bae 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsFifoIntegrationTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsFifoIntegrationTests.java @@ -93,6 +93,7 @@ * * @author Tomaz Fernandes * @author Mikhail Strokov + * @author José Iêdo */ @SpringBootTest class SqsFifoIntegrationTests extends BaseSqsIntegrationTest { @@ -123,6 +124,10 @@ class SqsFifoIntegrationTests extends BaseSqsIntegrationTest { static final String FIFO_VISIBILITY_TIMEOUT_EXTENSION_QUEUE_NAME = "fifo_visibility_timeout_extension_test_queue.fifo"; + static final String FIFO_SEND_MORE_THAN_10_SINGLE_GROUP_QUEUE_NAME = "fifo_send_more_than_10_single_group.fifo"; + + static final String FIFO_SEND_MORE_THAN_10_MULTIPLE_GROUPS_QUEUE_NAME = "fifo_send_more_than_10_multiple_groups.fifo"; + private static final String ERROR_ON_ACK_FACTORY = "errorOnAckFactory"; private static final String VISIBILITY_TIMEOUT_EXTENSION_FACTORY = "visibilityTimeoutExtensionFactory"; @@ -174,6 +179,8 @@ static void beforeTests() { createFifoQueue(client, FIFO_MANUALLY_CREATE_BATCH_CONTAINER_QUEUE_NAME), createFifoQueue(client, OBSERVES_MESSAGE_FIFO_QUEUE_NAME), createFifoQueue(client, FIFO_VISIBILITY_TIMEOUT_EXTENSION_QUEUE_NAME, getVisibilityAttribute("5")), + createFifoQueue(client, FIFO_SEND_MORE_THAN_10_SINGLE_GROUP_QUEUE_NAME), + createFifoQueue(client, FIFO_SEND_MORE_THAN_10_MULTIPLE_GROUPS_QUEUE_NAME), createFifoQueue(client, FIFO_MANUALLY_CREATE_BATCH_FACTORY_QUEUE_NAME)).join(); } @@ -535,6 +542,34 @@ void manuallyCreatesBatchFactory() throws Exception { assertThat(messagesContainer.manuallyCreatedBatchFactoryMessages).containsExactlyElementsOf(values); } + @Test + void shouldSendMoreThan10FifoMessagesInSingleGroup() { + String messageGroupId = UUID.randomUUID().toString(); + List> messages = IntStream.range(0, 25) + .mapToObj(i -> createMessage("payload-" + i, messageGroupId)).toList(); + SqsTemplate fifoTemplate = SqsTemplate.newTemplate(createAsyncClient()); + SendResult.Batch result = fifoTemplate.sendMany(FIFO_SEND_MORE_THAN_10_SINGLE_GROUP_QUEUE_NAME, + messages); + assertThat(result.successful()).hasSize(25); + assertThat(result.failed()).isEmpty(); + } + + @Test + void shouldSendMoreThan10FifoMessagesAcrossMultipleGroups() { + List valuesGroup1 = IntStream.range(0, 20).mapToObj(String::valueOf).collect(toList()); + List valuesGroup2 = IntStream.range(0, 15).mapToObj(String::valueOf).collect(toList()); + String group1 = UUID.randomUUID().toString(); + String group2 = UUID.randomUUID().toString(); + List> messages = new ArrayList<>(); + messages.addAll(createMessagesFromValues(group1, valuesGroup1)); + messages.addAll(createMessagesFromValues(group2, valuesGroup2)); + SqsTemplate fifoTemplate = SqsTemplate.newTemplate(createAsyncClient()); + SendResult.Batch result = fifoTemplate.sendMany(FIFO_SEND_MORE_THAN_10_MULTIPLE_GROUPS_QUEUE_NAME, + messages); + assertThat(result.successful()).hasSize(35); + assertThat(result.failed()).isEmpty(); + } + private Message createMessage(String body, String messageGroupId) { return MessageBuilder.withPayload(body) .setHeader(SqsHeaders.MessageSystemAttributes.SQS_MESSAGE_GROUP_ID_HEADER, messageGroupId) diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsIntegrationTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsIntegrationTests.java index 46e1f7b46..ac2a2f6d8 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsIntegrationTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsIntegrationTests.java @@ -145,6 +145,8 @@ class SqsIntegrationTests extends BaseSqsIntegrationTest { static final String MAX_CONCURRENT_MESSAGES_QUEUE_NAME = "max_concurrent_messages_test_queue"; + static final String SEND_MORE_THAN_10_MESSAGES_AT_ONCE_QUEUE_NAME = "send_more_than_10_message_test_queue"; + static final String LOW_RESOURCE_FACTORY = "lowResourceFactory"; static final String MANUAL_ACK_FACTORY = "manualAcknowledgementFactory"; @@ -178,7 +180,8 @@ static void beforeTests() { createQueue(client, MANUALLY_CREATE_FACTORY_QUEUE_NAME), createQueue(client, CONSUMES_ONE_MESSAGE_AT_A_TIME_QUEUE_NAME), createQueue(client, OBSERVES_MESSAGE_QUEUE_NAME), createQueue(client, OBSERVES_ERROR_QUEUE_NAME), - createQueue(client, MAX_CONCURRENT_MESSAGES_QUEUE_NAME)).join(); + createQueue(client, MAX_CONCURRENT_MESSAGES_QUEUE_NAME), + createQueue(client, SEND_MORE_THAN_10_MESSAGES_AT_ONCE_QUEUE_NAME)).join(); } @Autowired @@ -342,6 +345,15 @@ void receivesMessageBatch() throws Exception { assertThat(latchContainer.acknowledgementCallbackBatchLatch.await(10, TimeUnit.SECONDS)).isTrue(); } + @Test + void shouldSendMoreThan10MessagesAtOnce() { + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("moreThan10-payload-" + i).build()).toList(); + SendResult.Batch result = sqsTemplate.sendMany(SEND_MORE_THAN_10_MESSAGES_AT_ONCE_QUEUE_NAME, messages); + assertThat(result.successful()).hasSize(25); + assertThat(result.failed()).isEmpty(); + } + @Test void receivesMessageAsync() throws Exception { String messageBody = "receivesMessageAsync-payload"; diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/operations/SqsTemplateTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/operations/SqsTemplateTests.java index 296543b1a..dc8f13809 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/operations/SqsTemplateTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/operations/SqsTemplateTests.java @@ -29,6 +29,7 @@ import io.awspring.cloud.sqs.listener.SqsHeaders; import io.awspring.cloud.sqs.support.converter.ContextAwareMessagingMessageConverter; import java.time.Duration; +import java.util.ArrayList; import java.util.Collection; import java.util.Collections; import java.util.Iterator; @@ -38,7 +39,10 @@ import java.util.UUID; import java.util.concurrent.CompletableFuture; import java.util.concurrent.CompletionException; +import java.util.concurrent.atomic.AtomicInteger; import java.util.function.Consumer; +import java.util.stream.Collectors; +import java.util.stream.IntStream; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.mockito.ArgumentCaptor; @@ -573,6 +577,215 @@ void shouldThrowIfHasFailedMessagesInBatchByDefault() { }); } + @Test + void shouldPartitionMessagesIntoBatchesOf10() { + String queue = "test-queue"; + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("payload-" + i).build()).toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + List captured = new ArrayList<>(); + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + captured.add(request); + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(request.entries())).build()); + }); + + SqsOperations template = SqsTemplate.newSyncTemplate(mockClient); + SendResult.Batch result = template.sendMany(queue, messages); + + assertThat(captured).hasSize(3); + assertThat(captured.get(0).entries()).hasSize(10); + assertThat(captured.get(1).entries()).hasSize(10); + assertThat(captured.get(2).entries()).hasSize(5); + assertThat(result.successful()).hasSize(25); + assertThat(result.failed()).isEmpty(); + } + + @Test + void shouldStopSendingFifoBatchesAfterPartialFailure() { + String queue = "test-queue.fifo"; + String groupId = "test-group"; + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("payload-" + i) + .setHeader(SqsHeaders.MessageSystemAttributes.SQS_MESSAGE_GROUP_ID_HEADER, groupId).build()) + .toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + AtomicInteger callCount = new AtomicInteger(); + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + List entries = request.entries(); + if (callCount.incrementAndGet() == 2) { + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(entries.subList(0, 5))) + .failed(failedEntries(entries.subList(5, entries.size()))).build()); + } + return CompletableFuture + .completedFuture(SendMessageBatchResponse.builder().successful(successEntries(entries)).build()); + }); + + SqsTemplate template = SqsTemplate.builder().configure( + options -> options.sendBatchFailureHandlingStrategy(SendBatchFailureHandlingStrategy.DO_NOT_THROW)) + .sqsAsyncClient(mockClient).build(); + SendResult.Batch result = template.sendMany(queue, messages); + + then(mockClient).should(times(2)).sendMessageBatch(any(SendMessageBatchRequest.class)); + assertThat(result.successful()).hasSize(15); + assertThat(result.failed()).hasSize(10); + assertThat(result.failed().stream() + .filter(f -> f.errorMessage().equals("Skipped due to previous batch failure")).count()).isEqualTo(5); + } + + @Test + void shouldGroupMessagesByMessageGroupIdForFifoQueues() { + String queue = "test-queue.fifo"; + String groupA = "group-a"; + String groupB = "group-b"; + List> messages = IntStream.range(0, 27).mapToObj(i -> MessageBuilder.withPayload("payload-" + i) + .setHeader(SqsHeaders.MessageSystemAttributes.SQS_MESSAGE_GROUP_ID_HEADER, i < 15 ? groupA : groupB) + .build()).toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + List captured = new ArrayList<>(); + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + captured.add(request); + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(request.entries())).build()); + }); + + SqsOperations template = SqsTemplate.newSyncTemplate(mockClient); + SendResult.Batch result = template.sendMany(queue, messages); + + assertThat(result.successful()).hasSize(27); + assertThat(result.failed()).isEmpty(); + assertThat(captured).hasSize(4); + assertThat(captured).allSatisfy(request -> { + String firstGroupId = request.entries().get(0).messageGroupId(); + assertThat(request.entries()).allMatch(entry -> firstGroupId.equals(entry.messageGroupId())); + }); + } + + @Test + void shouldBinPackAutoGeneratedGroupIdMessagesForFifoQueues() { + String queue = "test-queue.fifo"; + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("payload-" + i).build()).toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + List captured = new ArrayList<>(); + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + captured.add(request); + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(request.entries())).build()); + }); + + SqsOperations template = SqsTemplate.newSyncTemplate(mockClient); + SendResult.Batch result = template.sendMany(queue, messages); + + assertThat(captured).hasSize(3); + assertThat(captured.get(0).entries()).hasSize(10); + assertThat(captured.get(1).entries()).hasSize(10); + assertThat(captured.get(2).entries()).hasSize(5); + assertThat(result.successful()).hasSize(25); + assertThat(result.failed()).isEmpty(); + assertThat(captured).allSatisfy(request -> assertThat(request.entries()) + .allSatisfy(entry -> assertThat(entry.messageGroupId()).isNotNull())); + } + + @Test + void shouldHandleBatchExceptionForStandardQueues() { + String queue = "test-queue"; + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("payload-" + i).build()).toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + if (request.entries().get(0).messageBody().equals("payload-10")) { + return CompletableFuture.failedFuture(new RuntimeException("test exception")); + } + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(request.entries())).build()); + }); + + SqsTemplate template = SqsTemplate.builder().configure( + options -> options.sendBatchFailureHandlingStrategy(SendBatchFailureHandlingStrategy.DO_NOT_THROW)) + .sqsAsyncClient(mockClient).build(); + SendResult.Batch result = template.sendMany(queue, messages); + + assertThat(result.successful().size() + result.failed().size()).isEqualTo(25); + assertThat(result.failed()).isNotEmpty(); + assertThat(result.failed()).allSatisfy(f -> assertThat(f.errorMessage()).isEqualTo("test exception")); + } + + @Test + void shouldHandleBatchExceptionForFifoQueues() { + String queue = "test-queue.fifo"; + String groupId = "test-group"; + List> messages = IntStream.range(0, 25) + .mapToObj(i -> MessageBuilder.withPayload("payload-" + i) + .setHeader(SqsHeaders.MessageSystemAttributes.SQS_MESSAGE_GROUP_ID_HEADER, groupId).build()) + .toList(); + + GetQueueUrlResponse urlResponse = GetQueueUrlResponse.builder().queueUrl(queue).build(); + given(mockClient.getQueueUrl(any(GetQueueUrlRequest.class))) + .willReturn(CompletableFuture.completedFuture(urlResponse)); + mockQueueAttributes(mockClient, Map.of()); + + AtomicInteger callCount = new AtomicInteger(); + given(mockClient.sendMessageBatch(any(SendMessageBatchRequest.class))).willAnswer(invocation -> { + SendMessageBatchRequest request = invocation.getArgument(0); + if (callCount.incrementAndGet() == 2) { + return CompletableFuture.failedFuture(new RuntimeException("test fifo exception")); + } + return CompletableFuture.completedFuture( + SendMessageBatchResponse.builder().successful(successEntries(request.entries())).build()); + }); + + SqsTemplate template = SqsTemplate.builder().configure( + options -> options.sendBatchFailureHandlingStrategy(SendBatchFailureHandlingStrategy.DO_NOT_THROW)) + .sqsAsyncClient(mockClient).build(); + SendResult.Batch result = template.sendMany(queue, messages); + + then(mockClient).should(times(2)).sendMessageBatch(any(SendMessageBatchRequest.class)); + assertThat(result.successful()).hasSize(10); + assertThat(result.failed()).hasSize(15); + assertThat(result.failed().stream().filter(f -> f.errorMessage().equals("test fifo exception")).count()) + .isEqualTo(10); + assertThat(result.failed().stream() + .filter(f -> f.errorMessage().equals("Skipped due to previous batch failure")).count()).isEqualTo(5); + } + + private static List sqsMessages(int count) { + return IntStream.range(0, count) + .mapToObj(i -> software.amazon.awssdk.services.sqs.model.Message.builder().build()) + .collect(Collectors.toCollection(ArrayList::new)); + } + @Test void shouldCreateByDefaultIfQueueNotFound() { String queue = "test-queue"; @@ -1371,4 +1584,17 @@ void shouldCacheSuccessfulQueueAttributesWithAttributeNames() { then(mockClient).should(times(1)).getQueueAttributes(any(Consumer.class)); then(mockClient).should(times(2)).sendMessage(any(SendMessageRequest.class)); } + + @SuppressWarnings("unchecked") + private static Consumer[] successEntries( + List entries) { + return entries.stream().> map( + e -> b -> b.id(e.id()).messageId(UUID.randomUUID().toString())).toArray(Consumer[]::new); + } + + @SuppressWarnings("unchecked") + private static Consumer[] failedEntries(List entries) { + return entries.stream().> map( + e -> b -> b.id(e.id()).message("error").code("ERR").senderFault(true)).toArray(Consumer[]::new); + } }