diff --git a/spring-cloud-stream-binder-kafka-core/src/main/java/org/springframework/cloud/stream/binder/kafka/properties/KafkaBinderConfigurationProperties.java b/spring-cloud-stream-binder-kafka-core/src/main/java/org/springframework/cloud/stream/binder/kafka/properties/KafkaBinderConfigurationProperties.java index 8cdd38154..2a0ab6201 100644 --- a/spring-cloud-stream-binder-kafka-core/src/main/java/org/springframework/cloud/stream/binder/kafka/properties/KafkaBinderConfigurationProperties.java +++ b/spring-cloud-stream-binder-kafka-core/src/main/java/org/springframework/cloud/stream/binder/kafka/properties/KafkaBinderConfigurationProperties.java @@ -27,6 +27,7 @@ import org.apache.kafka.clients.producer.ProducerConfig; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.autoconfigure.kafka.KafkaProperties; import org.springframework.boot.context.properties.ConfigurationProperties; +import org.springframework.kafka.support.DefaultKafkaHeaderMapper; import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; @@ -98,6 +99,11 @@ public class KafkaBinderConfigurationProperties { private JaasLoginModuleConfiguration jaas; + /** + * The bean name of a custom header mapper to use instead of a {@link DefaultKafkaHeaderMapper}. + */ + private String headerMapperBeanName; + public Transaction getTransaction() { return this.transaction; } @@ -358,6 +364,14 @@ public class KafkaBinderConfigurationProperties { this.jaas = jaas; } + public String getHeaderMapperBeanName() { + return this.headerMapperBeanName; + } + + public void setHeaderMapperBeanName(String headerMapperBeanName) { + this.headerMapperBeanName = headerMapperBeanName; + } + public static class Transaction { private final KafkaProducerProperties producer = new KafkaProducerProperties(); diff --git a/spring-cloud-stream-binder-kafka/src/main/java/org/springframework/cloud/stream/binder/kafka/KafkaMessageChannelBinder.java b/spring-cloud-stream-binder-kafka/src/main/java/org/springframework/cloud/stream/binder/kafka/KafkaMessageChannelBinder.java index ff21fb643..1185d0f3e 100644 --- a/spring-cloud-stream-binder-kafka/src/main/java/org/springframework/cloud/stream/binder/kafka/KafkaMessageChannelBinder.java +++ b/spring-cloud-stream-binder-kafka/src/main/java/org/springframework/cloud/stream/binder/kafka/KafkaMessageChannelBinder.java @@ -33,6 +33,7 @@ import org.apache.kafka.clients.producer.Producer; import org.apache.kafka.clients.producer.ProducerConfig; import org.apache.kafka.clients.producer.ProducerRecord; import org.apache.kafka.common.PartitionInfo; +import org.apache.kafka.common.header.Headers; import org.apache.kafka.common.serialization.ByteArrayDeserializer; import org.apache.kafka.common.serialization.ByteArraySerializer; import org.apache.kafka.common.utils.Utils; @@ -43,6 +44,7 @@ import org.springframework.cloud.stream.binder.BinderHeaders; import org.springframework.cloud.stream.binder.ExtendedConsumerProperties; import org.springframework.cloud.stream.binder.ExtendedProducerProperties; import org.springframework.cloud.stream.binder.ExtendedPropertiesBinder; +import org.springframework.cloud.stream.binder.HeaderMode; import org.springframework.cloud.stream.binder.kafka.properties.KafkaBinderConfigurationProperties; import org.springframework.cloud.stream.binder.kafka.properties.KafkaConsumerProperties; import org.springframework.cloud.stream.binder.kafka.properties.KafkaExtendedBindingProperties; @@ -67,6 +69,7 @@ import org.springframework.kafka.listener.AbstractMessageListenerContainer; import org.springframework.kafka.listener.ConcurrentMessageListenerContainer; import org.springframework.kafka.listener.config.ContainerProperties; import org.springframework.kafka.support.DefaultKafkaHeaderMapper; +import org.springframework.kafka.support.KafkaHeaderMapper; import org.springframework.kafka.support.KafkaHeaders; import org.springframework.kafka.support.ProducerListener; import org.springframework.kafka.support.SendResult; @@ -112,7 +115,7 @@ public class KafkaMessageChannelBinder extends public KafkaMessageChannelBinder(KafkaBinderConfigurationProperties configurationProperties, KafkaTopicProvisioner provisioningProvider) { - super(true, null, provisioningProvider); + super(headersToMap(configurationProperties), provisioningProvider); this.configurationProperties = configurationProperties; if (StringUtils.hasText(configurationProperties.getTransaction().getTransactionIdPrefix())) { this.transactionManager = new KafkaTransactionManager<>( @@ -124,6 +127,22 @@ public class KafkaMessageChannelBinder extends } } + private static String[] headersToMap(KafkaBinderConfigurationProperties configurationProperties) { + String[] headersToMap; + if (ObjectUtils.isEmpty(configurationProperties.getHeaders())) { + headersToMap = BinderHeaders.STANDARD_HEADERS; + } + else { + String[] combinedHeadersToMap = Arrays.copyOfRange(BinderHeaders.STANDARD_HEADERS, 0, + BinderHeaders.STANDARD_HEADERS.length + configurationProperties.getHeaders().length); + System.arraycopy(configurationProperties.getHeaders(), 0, combinedHeadersToMap, + BinderHeaders.STANDARD_HEADERS.length, + configurationProperties.getHeaders().length); + headersToMap = combinedHeadersToMap; + } + return headersToMap; + } + public void setExtendedBindingProperties(KafkaExtendedBindingProperties extendedBindingProperties) { this.extendedBindingProperties = extendedBindingProperties; } @@ -194,19 +213,37 @@ public class KafkaMessageChannelBinder extends if (errorChannel != null) { handler.setSendFailureChannel(errorChannel); } - String[] headerPatterns = producerProperties.getExtension().getHeaderPatterns(); - if (headerPatterns != null && headerPatterns.length > 0) { - List patterns = new LinkedList<>(Arrays.asList(headerPatterns)); - if (!patterns.contains("!" + MessageHeaders.TIMESTAMP)) { - patterns.add(0, "!" + MessageHeaders.TIMESTAMP); - } - if (!patterns.contains("!" + MessageHeaders.ID)) { - patterns.add(0, "!" + MessageHeaders.ID); - } - DefaultKafkaHeaderMapper headerMapper = new DefaultKafkaHeaderMapper( - patterns.toArray(new String[patterns.size()])); - handler.setHeaderMapper(headerMapper); + KafkaHeaderMapper mapper = null; + if (this.configurationProperties.getHeaderMapperBeanName() != null) { + mapper = getApplicationContext().getBean(this.configurationProperties.getHeaderMapperBeanName(), + KafkaHeaderMapper.class); } + /* + * Even if the user configures a bean, we must not use it if the header + * mode is not the default (headers); setting the mapper to null + * disables populating headers in the message handler. + */ + if (producerProperties.getHeaderMode() != null + && !HeaderMode.headers.equals(producerProperties.getHeaderMode())) { + mapper = null; + } + else if (mapper == null) { + String[] headerPatterns = producerProperties.getExtension().getHeaderPatterns(); + if (headerPatterns != null && headerPatterns.length > 0) { + List patterns = new LinkedList<>(Arrays.asList(headerPatterns)); + if (!patterns.contains("!" + MessageHeaders.TIMESTAMP)) { + patterns.add(0, "!" + MessageHeaders.TIMESTAMP); + } + if (!patterns.contains("!" + MessageHeaders.ID)) { + patterns.add(0, "!" + MessageHeaders.ID); + } + mapper = new DefaultKafkaHeaderMapper(patterns.toArray(new String[patterns.size()])); + } + else { + mapper = new DefaultKafkaHeaderMapper(); + } + } + handler.setHeaderMapper(mapper); return handler; } @@ -326,12 +363,30 @@ public class KafkaMessageChannelBinder extends final KafkaMessageDrivenChannelAdapter kafkaMessageDrivenChannelAdapter = new KafkaMessageDrivenChannelAdapter<>( messageListenerContainer); MessagingMessageConverter messageConverter = new MessagingMessageConverter(); - DefaultKafkaHeaderMapper headerMapper = new DefaultKafkaHeaderMapper(); - String[] trustedPackages = extendedConsumerProperties.getExtension().getTrustedPackages(); - if (!StringUtils.isEmpty(trustedPackages)) { - headerMapper.addTrustedPackages(trustedPackages); + KafkaHeaderMapper mapper = null; + if (this.configurationProperties.getHeaderMapperBeanName() != null) { + mapper = getApplicationContext().getBean(this.configurationProperties.getHeaderMapperBeanName(), + KafkaHeaderMapper.class); } - messageConverter.setHeaderMapper(headerMapper); + if (mapper == null) { + DefaultKafkaHeaderMapper headerMapper = new DefaultKafkaHeaderMapper() { + + @Override + public void toHeaders(Headers source, Map headers) { + super.toHeaders(source, headers); + if (headers.size() > 0) { + headers.put(BinderHeaders.NATIVE_HEADERS_PRESENT, Boolean.TRUE); + } + } + + }; + String[] trustedPackages = extendedConsumerProperties.getExtension().getTrustedPackages(); + if (!StringUtils.isEmpty(trustedPackages)) { + headerMapper.addTrustedPackages(trustedPackages); + } + mapper = headerMapper; + } + messageConverter.setHeaderMapper(mapper); kafkaMessageDrivenChannelAdapter.setMessageConverter(messageConverter); kafkaMessageDrivenChannelAdapter.setBeanFactory(this.getBeanFactory()); ErrorInfrastructure errorInfrastructure = registerErrorInfrastructure(destination, consumerGroup, diff --git a/spring-cloud-stream-binder-kafka/src/test/java/org/springframework/cloud/stream/binder/kafka/KafkaBinderTests.java b/spring-cloud-stream-binder-kafka/src/test/java/org/springframework/cloud/stream/binder/kafka/KafkaBinderTests.java index c2c7edbdc..772605046 100644 --- a/spring-cloud-stream-binder-kafka/src/test/java/org/springframework/cloud/stream/binder/kafka/KafkaBinderTests.java +++ b/spring-cloud-stream-binder-kafka/src/test/java/org/springframework/cloud/stream/binder/kafka/KafkaBinderTests.java @@ -19,7 +19,9 @@ package org.springframework.cloud.stream.binder.kafka; import java.io.IOException; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; import java.util.HashMap; +import java.util.Iterator; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; @@ -30,12 +32,11 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; -import com.fasterxml.jackson.databind.ObjectMapper; -import kafka.utils.ZKStringSerializer$; -import kafka.utils.ZkUtils; - import org.I0Itec.zkclient.ZkClient; +import org.apache.kafka.clients.consumer.Consumer; import org.apache.kafka.clients.consumer.ConsumerConfig; +import org.apache.kafka.clients.consumer.ConsumerRecord; +import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.producer.ProducerConfig; import org.apache.kafka.clients.producer.ProducerRecord; @@ -91,6 +92,7 @@ import org.springframework.kafka.support.SendResult; import org.springframework.kafka.support.TopicPartitionInitialOffset; import org.springframework.kafka.test.core.BrokerAddress; import org.springframework.kafka.test.rule.KafkaEmbedded; +import org.springframework.kafka.test.utils.KafkaTestUtils; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; import org.springframework.messaging.MessageHandler; @@ -108,11 +110,16 @@ import org.springframework.util.MimeTypeUtils; import org.springframework.util.concurrent.ListenableFuture; import org.springframework.util.concurrent.SettableListenableFuture; +import com.fasterxml.jackson.databind.ObjectMapper; + import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.fail; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; +import kafka.utils.ZKStringSerializer$; +import kafka.utils.ZkUtils; + /** * @author Soby Chacko @@ -1759,7 +1766,7 @@ public class KafkaBinderTests extends public void testPartitionedModuleJavaWithRawMode() throws Exception { Binder binder = getBinder(); ExtendedProducerProperties properties = createProducerProperties(); - properties.setHeaderMode(HeaderMode.raw); + properties.setHeaderMode(HeaderMode.none); properties.setPartitionKeyExtractorClass(RawKafkaPartitionTestSupport.class); properties.setPartitionSelectorClass(RawKafkaPartitionTestSupport.class); properties.setPartitionCount(6); @@ -1773,7 +1780,7 @@ public class KafkaBinderTests extends consumerProperties.setInstanceCount(3); consumerProperties.setInstanceIndex(0); consumerProperties.setPartitioned(true); - consumerProperties.setHeaderMode(HeaderMode.raw); + consumerProperties.setHeaderMode(HeaderMode.none); consumerProperties.getExtension().setAutoRebalanceEnabled(false); QueueChannel input0 = new QueueChannel(); input0.setBeanName("test.input0J"); @@ -1815,7 +1822,7 @@ public class KafkaBinderTests extends properties.setPartitionKeyExpression(spelExpressionParser.parseExpression("payload[0]")); properties.setPartitionSelectorExpression(spelExpressionParser.parseExpression("hashCode()")); properties.setPartitionCount(6); - properties.setHeaderMode(HeaderMode.raw); + properties.setHeaderMode(HeaderMode.none); DirectChannel output = createBindableChannel("output", createProducerBindingProperties(properties)); output.setBeanName("test.output"); @@ -1833,7 +1840,7 @@ public class KafkaBinderTests extends consumerProperties.setInstanceIndex(0); consumerProperties.setInstanceCount(3); consumerProperties.setPartitioned(true); - consumerProperties.setHeaderMode(HeaderMode.raw); + consumerProperties.setHeaderMode(HeaderMode.none); consumerProperties.getExtension().setAutoRebalanceEnabled(false); QueueChannel input0 = new QueueChannel(); input0.setBeanName("test.input0S"); @@ -1875,11 +1882,11 @@ public class KafkaBinderTests extends DirectChannel moduleOutputChannel = new DirectChannel(); QueueChannel moduleInputChannel = new QueueChannel(); ExtendedProducerProperties producerProperties = createProducerProperties(); - producerProperties.setHeaderMode(HeaderMode.raw); + producerProperties.setHeaderMode(HeaderMode.none); Binding producerBinding = binder.bindProducer("raw.0", moduleOutputChannel, producerProperties); ExtendedConsumerProperties consumerProperties = createConsumerProperties(); - consumerProperties.setHeaderMode(HeaderMode.raw); + consumerProperties.setHeaderMode(HeaderMode.none); Binding consumerBinding = binder.bindConsumer("raw.0", "test", moduleInputChannel, consumerProperties); Message message = org.springframework.integration.support.MessageBuilder @@ -1894,6 +1901,95 @@ public class KafkaBinderTests extends consumerBinding.unbind(); } + /* + * Verify that a consumer configured to handle embedded headers can handle + * all three variants. + */ + @Test + @SuppressWarnings({ "unchecked", "rawtypes" }) + public void testSendAndReceiveWithMixedMode() throws Exception { + KafkaBinderConfigurationProperties binderConfiguration = createConfigurationProperties(); + binderConfiguration.setHeaders("foo"); + Binder binder = getBinder(binderConfiguration); + QueueChannel moduleInputChannel = new QueueChannel(); + DirectChannel moduleOutputChannel1 = new DirectChannel(); + ExtendedProducerProperties producerProperties1 = createProducerProperties(); + producerProperties1.setHeaderMode(HeaderMode.embeddedHeaders); + Binding producerBinding1 = binder.bindProducer("mixed.0", moduleOutputChannel1, + producerProperties1); + + DirectChannel moduleOutputChannel2 = new DirectChannel(); + ExtendedProducerProperties producerProperties2 = createProducerProperties(); + producerProperties2.setHeaderMode(HeaderMode.headers); + Binding producerBinding2 = binder.bindProducer("mixed.0", moduleOutputChannel2, + producerProperties2); + + DirectChannel moduleOutputChannel3 = new DirectChannel(); + ExtendedProducerProperties producerProperties3 = createProducerProperties(); + producerProperties3.setHeaderMode(HeaderMode.none); + Binding producerBinding3 = binder.bindProducer("mixed.0", moduleOutputChannel3, + producerProperties3); + + ExtendedConsumerProperties consumerProperties = createConsumerProperties(); + consumerProperties.setHeaderMode(HeaderMode.embeddedHeaders); + Binding consumerBinding = binder.bindConsumer("mixed.0", "test", moduleInputChannel, + consumerProperties); + Message message = org.springframework.integration.support.MessageBuilder + .withPayload("testSendAndReceiveWithMixedMode".getBytes()) + .setHeader("foo", "bar") + .build(); + // Let the consumer actually bind to the producer before sending a msg + binderBindUnbindLatency(); + moduleOutputChannel1.send(message); + moduleOutputChannel2.send(message); + moduleOutputChannel3.send(message); + Message inbound = receive(moduleInputChannel, 10_000); + assertThat(inbound).isNotNull(); + assertThat(new String((byte[]) inbound.getPayload())).isEqualTo("testSendAndReceiveWithMixedMode"); + assertThat(inbound.getHeaders().get("foo")).isEqualTo("bar"); + assertThat(inbound.getHeaders().get(BinderHeaders.NATIVE_HEADERS_PRESENT)).isNull(); + inbound = receive(moduleInputChannel); + assertThat(inbound).isNotNull(); + assertThat(new String((byte[]) inbound.getPayload())).isEqualTo("testSendAndReceiveWithMixedMode"); + assertThat(inbound.getHeaders().get("foo")).isEqualTo("bar"); + assertThat(inbound.getHeaders().get(BinderHeaders.NATIVE_HEADERS_PRESENT)).isEqualTo(Boolean.TRUE); + inbound = receive(moduleInputChannel); + assertThat(inbound).isNotNull(); + assertThat(new String((byte[]) inbound.getPayload())).isEqualTo("testSendAndReceiveWithMixedMode"); + assertThat(inbound.getHeaders().get("foo")).isNull(); + assertThat(inbound.getHeaders().get(BinderHeaders.NATIVE_HEADERS_PRESENT)).isNull(); + + Map consumerProps = KafkaTestUtils.consumerProps("testSendAndReceiveWithMixedMode", "false", + embeddedKafka); + consumerProps.put(ConsumerConfig.AUTO_OFFSET_RESET_CONFIG, "earliest"); + consumerProps.put(ConsumerConfig.KEY_DESERIALIZER_CLASS_CONFIG, ByteArrayDeserializer.class); + consumerProps.put(ConsumerConfig.VALUE_DESERIALIZER_CLASS_CONFIG, ByteArrayDeserializer.class); + DefaultKafkaConsumerFactory cf = new DefaultKafkaConsumerFactory<>(consumerProps); + Consumer consumer = cf.createConsumer(); + consumer.subscribe(Collections.singletonList("mixed.0")); + + ConsumerRecords records = consumer.poll(10_1000); + Iterator iterator = records.iterator(); + ConsumerRecord record = iterator.next(); + byte[] value = (byte[]) record.value(); + assertThat(value[0] & 0xff).isEqualTo(0xff); + assertThat(record.headers().toArray().length).isEqualTo(0); + record = iterator.next(); + value = (byte[]) record.value(); + assertThat(value[0] & 0xff).isNotEqualTo(0xff); + assertThat(record.headers().toArray().length).isEqualTo(2); + record = iterator.next(); + value = (byte[]) record.value(); + assertThat(value[0] & 0xff).isNotEqualTo(0xff); + assertThat(record.headers().toArray().length).isEqualTo(0); + consumer.close(); + + producerBinding1.unbind(); + producerBinding2.unbind(); + producerBinding3.unbind(); + consumerBinding.unbind(); + } + @SuppressWarnings({ "rawtypes", "unchecked" }) @Test public void testProducerErrorChannel() throws Exception { @@ -1901,7 +1997,7 @@ public class KafkaBinderTests extends DirectChannel moduleOutputChannel = createBindableChannel("output", new BindingProperties()); ExtendedProducerProperties producerProps = new ExtendedProducerProperties<>( new KafkaProducerProperties()); - producerProps.setHeaderMode(HeaderMode.raw); + producerProps.setHeaderMode(HeaderMode.none); producerProps.setErrorChannelEnabled(true); Binding producerBinding = binder.bindProducer("ec.0", moduleOutputChannel, producerProps); final Message message = MessageBuilder.withPayload("bad").setHeader(MessageHeaders.CONTENT_TYPE, "application/json")