diff --git a/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/binder/BinderHeaders.java b/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/binder/BinderHeaders.java index 515b4147a..091a19c2f 100644 --- a/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/binder/BinderHeaders.java +++ b/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/binder/BinderHeaders.java @@ -39,6 +39,12 @@ public final class BinderHeaders { IntegrationMessageHeaderAccessor.SEQUENCE_NUMBER, MessageHeaders.CONTENT_TYPE}; private static final String PREFIX = "scst_"; + + + /** + * Name of the Message header identifying structure for batch Message headers. + */ + public static String BATCH_HEADERS = PREFIX + "batchHeaders"; /** * Indicates the name of the target destination the binder should use if they diff --git a/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/function/StandardBatchUtils.java b/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/function/StandardBatchUtils.java index 178f2932d..3f23cb42a 100644 --- a/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/function/StandardBatchUtils.java +++ b/core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/function/StandardBatchUtils.java @@ -18,21 +18,81 @@ package org.springframework.cloud.stream.function; import java.util.ArrayList; import java.util.HashMap; +import java.util.Iterator; import java.util.List; import java.util.Map; +import java.util.Map.Entry; +import org.springframework.cloud.stream.binder.BinderHeaders; import org.springframework.messaging.Message; import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.support.MessageBuilder; +import org.springframework.util.Assert; /** * @author Oleg Zhurakousky * @since 4.2 */ -public class StandardBatchUtils { +public final class StandardBatchUtils { - public static String BATCH_HEADERS = "scst_batchHeaders"; + private StandardBatchUtils() { + + } + /** + * Iterates over batch message structure returning {@link Iterable} of individual messages. + * + * @param batchMessage instance of batch {@link Message} + * @return instance of {@link Iterable} representing individual Messages in a batch {@link Message} as {@link Entry}. + */ + public static Iterable>> iterate(Message> batchMessage) { + return new Iterable>>() { + @Override + public Iterator>> iterator() { + return new Iterator>>() { + int index = 0; + @Override + public Entry> next() { + return getMessageByIndex(batchMessage, index++); + } + + @Override + public boolean hasNext() { + return index < batchMessage.getPayload().size(); + } + }; + } + }; + } + + /** + * Extracts individual {@link Message} by index from batch {@link Message} + * @param batchMessage instance of batch {@link Message} + * @param index index of individual {@link Message} in a batch + * @return individual {@link Message} in a batch {@link Message} + */ + public static Entry> getMessageByIndex(Message> batchMessage, int index) { + Assert.isTrue(index < batchMessage.getPayload().size(), "Index " + index + " is out of bounds as there are only " + + batchMessage.getPayload().size() + " messages in a batch."); + return new Entry>() { + + @Override + public Map setValue(Map value) { + throw new UnsupportedOperationException(); + } + + @SuppressWarnings("unchecked") + @Override + public Map getValue() { + return ((List>) batchMessage.getHeaders().get(BinderHeaders.BATCH_HEADERS)).get(index); + } + + @Override + public Object getKey() { + return batchMessage.getPayload().get(index); + } + }; + } public static class BatchMessageBuilder { @@ -48,13 +108,13 @@ public class StandardBatchUtils { return this; } - public BatchMessageBuilder addHeader(String key, Object value) { + public BatchMessageBuilder addRootHeader(String key, Object value) { this.headers.put(key, value); return this; } public Message> build() { - this.headers.put(BATCH_HEADERS, this.batchHeaders); + this.headers.put(BinderHeaders.BATCH_HEADERS, this.batchHeaders); return MessageBuilder.createMessage(payloads, new MessageHeaders(headers)); } } diff --git a/core/spring-cloud-stream/src/test/java/org/springframework/cloud/stream/function/StandardBatchUtilsTests.java b/core/spring-cloud-stream/src/test/java/org/springframework/cloud/stream/function/StandardBatchUtilsTests.java index 77c387cfa..936efd0a5 100644 --- a/core/spring-cloud-stream/src/test/java/org/springframework/cloud/stream/function/StandardBatchUtilsTests.java +++ b/core/spring-cloud-stream/src/test/java/org/springframework/cloud/stream/function/StandardBatchUtilsTests.java @@ -18,14 +18,18 @@ package org.springframework.cloud.stream.function; import static org.assertj.core.api.Assertions.assertThat; +import java.util.ArrayList; import java.util.Collections; import java.util.List; import java.util.Map; +import java.util.Map.Entry; import org.junit.jupiter.api.Test; +import org.springframework.cloud.stream.binder.BinderHeaders; import org.springframework.cloud.stream.function.StandardBatchUtils.BatchMessageBuilder; import org.springframework.messaging.Message; + /** * */ @@ -36,7 +40,7 @@ public class StandardBatchUtilsTests { public void testBatchMessageBuilder() { BatchMessageBuilder builder = new BatchMessageBuilder(); builder.addMessage("foo", Collections.singletonMap("fooKey", "fooValue")); - builder.addHeader("a", "a"); + builder.addRootHeader("a", "a"); builder.addMessage("bar", Collections.singletonMap("barKey", "barValue")); builder.addMessage("baz", Collections.singletonMap("bazKey", "bazValue")); @@ -45,7 +49,7 @@ public class StandardBatchUtilsTests { List payloads = batchMessage.getPayload(); assertThat(payloads.size()).isEqualTo(3); - List> batchHeaders = (List>) batchMessage.getHeaders().get(StandardBatchUtils.BATCH_HEADERS); + List> batchHeaders = (List>) batchMessage.getHeaders().get(BinderHeaders.BATCH_HEADERS); assertThat(batchHeaders.size()).isEqualTo(3); assertThat(payloads.get(0)).isEqualTo("foo"); @@ -56,4 +60,29 @@ public class StandardBatchUtilsTests { assertThat(batchMessage.getHeaders().get("a")).isEqualTo("a"); } + + @Test + public void testIterator() { + BatchMessageBuilder builder = new BatchMessageBuilder(); + builder.addMessage("foo", Collections.singletonMap("fooKey", "fooValue")); + builder.addRootHeader("a", "a"); + builder.addMessage("bar", Collections.singletonMap("barKey", "barValue")); + builder.addMessage("baz", Collections.singletonMap("bazKey", "bazValue")); + + Message> batchMessage = builder.build(); + + List>> entries = new ArrayList<>(); + StandardBatchUtils.iterate(batchMessage).forEach(entry -> { + entries.add(entry); + }); + assertThat(entries.size()).isEqualTo(3); + assertThat(entries.get(0).getKey()).isEqualTo("foo"); + assertThat(entries.get(0).getValue().get("fooKey")).isEqualTo("fooValue"); + + assertThat(entries.get(1).getKey()).isEqualTo("bar"); + assertThat(entries.get(1).getValue().get("barKey")).isEqualTo("barValue"); + + assertThat(entries.get(2).getKey()).isEqualTo("baz"); + assertThat(entries.get(2).getValue().get("bazKey")).isEqualTo("bazValue"); + } }