diff --git a/spring-integration-amqp/src/main/java/org/springframework/integration/amqp/inbound/AmqpInboundChannelAdapter.java b/spring-integration-amqp/src/main/java/org/springframework/integration/amqp/inbound/AmqpInboundChannelAdapter.java index e5410846b7..7f8a9c1ef9 100644 --- a/spring-integration-amqp/src/main/java/org/springframework/integration/amqp/inbound/AmqpInboundChannelAdapter.java +++ b/spring-integration-amqp/src/main/java/org/springframework/integration/amqp/inbound/AmqpInboundChannelAdapter.java @@ -124,6 +124,8 @@ public class AmqpInboundChannelAdapter extends MessageProducerSupport implements private BatchMode batchMode = BatchMode.MESSAGES; + private String headerNameForBatchedHeaders = CONSOLIDATED_HEADERS; + /** * Construct an instance using the provided container. * @param listenerContainer the container. @@ -137,7 +139,8 @@ public class AmqpInboundChannelAdapter extends MessageProducerSupport implements this.messageListenerContainer = listenerContainer; this.messageListenerContainer.setAutoStartup(false); setErrorMessageStrategy(new AmqpMessageHeaderErrorMessageStrategy()); - this.abstractListenerContainer = listenerContainer instanceof AbstractMessageListenerContainer abstractMessageListenerContainer + this.abstractListenerContainer = + listenerContainer instanceof AbstractMessageListenerContainer abstractMessageListenerContainer ? abstractMessageListenerContainer : null; } @@ -220,6 +223,20 @@ public class AmqpInboundChannelAdapter extends MessageProducerSupport implements this.batchMode = batchMode; } + /** + * Set a header name containing {@code List} headers when batch mode + * is {@link BatchMode#EXTRACT_PAYLOADS_WITH_HEADERS}. + * Defaults to {@link #CONSOLIDATED_HEADERS}. + * @param headerNameForBatchedHeaders the name of header + * containing {@code List} headers when batch mode + * is {@link BatchMode#EXTRACT_PAYLOADS_WITH_HEADERS}. + * @since 6.4 + */ + public void setHeaderNameForBatchedHeaders(String headerNameForBatchedHeaders) { + Assert.hasText(headerNameForBatchedHeaders, "'headerNameForBatchedHeaders' must not be empty"); + this.headerNameForBatchedHeaders = headerNameForBatchedHeaders; + } + @Override public String getComponentType() { return "amqp:inbound-channel-adapter"; @@ -436,7 +453,7 @@ public class AmqpInboundChannelAdapter extends MessageProducerSupport implements headers.put(IntegrationMessageHeaderAccessor.DELIVERY_ATTEMPT, new AtomicInteger()); } if (listHeaders != null) { - headers.put(CONSOLIDATED_HEADERS, listHeaders); + headers.put(AmqpInboundChannelAdapter.this.headerNameForBatchedHeaders, listHeaders); } return getMessageBuilderFactory() .withPayload(payload) diff --git a/spring-integration-amqp/src/test/java/org/springframework/integration/amqp/inbound/InboundEndpointTests.java b/spring-integration-amqp/src/test/java/org/springframework/integration/amqp/inbound/InboundEndpointTests.java index 5cbb987910..06184fa541 100644 --- a/spring-integration-amqp/src/test/java/org/springframework/integration/amqp/inbound/InboundEndpointTests.java +++ b/spring-integration-amqp/src/test/java/org/springframework/integration/amqp/inbound/InboundEndpointTests.java @@ -75,7 +75,7 @@ import static org.mockito.ArgumentMatchers.anyBoolean; import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.isNull; import static org.mockito.BDDMockito.given; -import static org.mockito.Mockito.doAnswer; +import static org.mockito.BDDMockito.willReturn; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.spy; import static org.mockito.Mockito.when; @@ -90,9 +90,9 @@ public class InboundEndpointTests { @Test public void testInt2809JavaTypePropertiesToAmqp() throws Exception { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -104,7 +104,7 @@ public class InboundEndpointTests { PollableChannel channel = new QueueChannel(); adapter.setOutputChannel(channel); - adapter.setBeanFactory(mock(BeanFactory.class)); + adapter.setBeanFactory(mock()); adapter.setBindSourceMessage(true); adapter.afterPropertiesSet(); @@ -134,9 +134,9 @@ public class InboundEndpointTests { @Test public void testInt2809JavaTypePropertiesFromAmqp() throws Exception { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -169,9 +169,9 @@ public class InboundEndpointTests { @Test public void testMessageConverterJsonHeadersHavePrecedenceOverMessageHeaders() throws Exception { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -235,9 +235,9 @@ public class InboundEndpointTests { @Test public void testAdapterConversionError() throws Exception { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -285,9 +285,9 @@ public class InboundEndpointTests { @Test public void testGatewayConversionError() throws Exception { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -339,7 +339,7 @@ public class InboundEndpointTests { @Test public void testRetryWithinOnMessageAdapter() throws Exception { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); AbstractMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); adapter.setOutputChannel(new DirectChannel()); @@ -367,7 +367,7 @@ public class InboundEndpointTests { @Test public void testRetryWithMessageRecovererOnMessageAdapter() throws Exception { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); AbstractMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); adapter.setOutputChannel(new DirectChannel()); @@ -400,7 +400,7 @@ public class InboundEndpointTests { @Test public void testRetryWithinOnMessageGateway() throws Exception { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); AbstractMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); AmqpInboundGateway adapter = new AmqpInboundGateway(container); adapter.setRequestChannel(new DirectChannel()); @@ -428,7 +428,7 @@ public class InboundEndpointTests { @Test public void testRetryWithMessageRecovererOnMessageGateway() throws Exception { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); AbstractMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); AmqpInboundGateway adapter = new AmqpInboundGateway(container); adapter.setRequestChannel(new DirectChannel()); @@ -462,7 +462,7 @@ public class InboundEndpointTests { @SuppressWarnings({"unchecked"}) @Test public void testBatchAdapter() throws Exception { - SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock(ConnectionFactory.class)); + SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock()); container.setDeBatchingEnabled(false); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); QueueChannel out = new QueueChannel(); @@ -486,7 +486,7 @@ public class InboundEndpointTests { @SuppressWarnings({"unchecked"}) @Test public void testBatchGateway() throws Exception { - SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock(ConnectionFactory.class)); + SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock()); container.setDeBatchingEnabled(false); AmqpInboundGateway gateway = new AmqpInboundGateway(container); QueueChannel out = new QueueChannel(); @@ -514,12 +514,13 @@ public class InboundEndpointTests { @SuppressWarnings({"unchecked"}) @Test public void testConsumerBatchExtract() { - SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock(ConnectionFactory.class)); + SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock()); container.setConsumerBatchEnabled(true); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); QueueChannel out = new QueueChannel(); adapter.setOutputChannel(out); adapter.setBatchMode(BatchMode.EXTRACT_PAYLOADS_WITH_HEADERS); + adapter.setHeaderNameForBatchedHeaders("some_batch_headers"); adapter.afterPropertiesSet(); ChannelAwareBatchMessageListener listener = (ChannelAwareBatchMessageListener) container.getMessageListener(); MessageProperties messageProperties = new MessageProperties(); @@ -531,14 +532,14 @@ public class InboundEndpointTests { Message received = out.receive(0); assertThat(received).isNotNull(); assertThat(((List) received.getPayload())).contains("test1", "test2"); - assertThat(received.getHeaders().get(AmqpInboundChannelAdapter.CONSOLIDATED_HEADERS, List.class)) + assertThat(received.getHeaders().get("some_batch_headers", List.class)) .hasSize(2); } @SuppressWarnings({"unchecked"}) @Test public void testConsumerBatch() { - SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock(ConnectionFactory.class)); + SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock()); container.setConsumerBatchEnabled(true); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); QueueChannel out = new QueueChannel(); @@ -560,7 +561,7 @@ public class InboundEndpointTests { @Test public void testConsumerBatchAndWrongMessageRecoverer() { - SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock(ConnectionFactory.class)); + SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(mock()); container.setConsumerBatchEnabled(true); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); adapter.setRetryTemplate(new RetryTemplate()); @@ -574,7 +575,7 @@ public class InboundEndpointTests { @Test public void testExclusiveRecover() { - AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(mock(AbstractMessageListenerContainer.class)); + AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(mock()); adapter.setRetryTemplate(new RetryTemplate()); adapter.setMessageRecoverer((message, cause) -> { }); @@ -587,9 +588,9 @@ public class InboundEndpointTests { @Test public void testAdapterConversionErrorConsumerBatchExtract() { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -644,9 +645,9 @@ public class InboundEndpointTests { @Test public void testAdapterConversionErrorConsumerBatch() { - Connection connection = mock(Connection.class); - doAnswer(invocation -> mock(Channel.class)).when(connection).createChannel(anyBoolean()); - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + Connection connection = mock(); + willReturn(mock(Channel.class)).given(connection).createChannel(anyBoolean()); + ConnectionFactory connectionFactory = mock(); when(connectionFactory.createConnection()).thenReturn(connection); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(); container.setConnectionFactory(connectionFactory); @@ -700,7 +701,7 @@ public class InboundEndpointTests { @Test public void testRetryWithinOnMessageAdapterConsumerBatch() { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); container.setConsumerBatchEnabled(true); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container); @@ -746,7 +747,7 @@ public class InboundEndpointTests { @Test public void testRetryWithMessageRecovererOnMessageAdapterConsumerBatch() throws InterruptedException { - ConnectionFactory connectionFactory = mock(ConnectionFactory.class); + ConnectionFactory connectionFactory = mock(); SimpleMessageListenerContainer container = new SimpleMessageListenerContainer(connectionFactory); container.setConsumerBatchEnabled(true); AmqpInboundChannelAdapter adapter = new AmqpInboundChannelAdapter(container);