GH-9647: Add AmqpInboundChannelAdapter.setHeaderNameForBatchedHeaders()

Fixes: #9647
Issue link: https://github.com/spring-projects/spring-integration/issues/9647
This commit is contained in:
Artem Bilan
2024-11-12 10:52:42 -05:00
parent 29cd21624c
commit a96b5f0a6b
2 changed files with 56 additions and 38 deletions

View File

@@ -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<Map<String, Object>} 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<Map<String, Object>} 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)

View File

@@ -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<String>) 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);