diff --git a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandler.java b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandler.java index db44cca239..7478f8ec5e 100644 --- a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandler.java +++ b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandler.java @@ -19,6 +19,10 @@ package org.springframework.integration.kafka.outbound; import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeoutException; +import org.apache.kafka.clients.producer.ProducerRecord; +import org.apache.kafka.common.header.Headers; +import org.apache.kafka.common.header.internals.RecordHeaders; + import org.springframework.expression.EvaluationContext; import org.springframework.expression.Expression; import org.springframework.integration.MessageTimeoutException; @@ -26,6 +30,9 @@ import org.springframework.integration.expression.ExpressionUtils; import org.springframework.integration.expression.ValueExpression; import org.springframework.integration.handler.AbstractMessageHandler; import org.springframework.kafka.core.KafkaTemplate; +import org.springframework.kafka.support.DefaultKafkaHeaderMapper; +import org.springframework.kafka.support.JacksonPresent; +import org.springframework.kafka.support.KafkaHeaderMapper; import org.springframework.kafka.support.KafkaHeaders; import org.springframework.kafka.support.KafkaNull; import org.springframework.messaging.Message; @@ -67,9 +74,14 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { private Expression sendTimeoutExpression = new ValueExpression<>(DEFAULT_SEND_TIMEOUT); + private KafkaHeaderMapper headerMapper; + public KafkaProducerMessageHandler(final KafkaTemplate kafkaTemplate) { Assert.notNull(kafkaTemplate, "kafkaTemplate cannot be null"); this.kafkaTemplate = kafkaTemplate; + if (JacksonPresent.isJackson2Present()) { + this.headerMapper = new DefaultKafkaHeaderMapper(); + } } public void setTopicExpression(Expression topicExpression) { @@ -90,12 +102,21 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { * * @param timestampExpression the {@link Expression} for timestamp to wait for result * fo send operation. - * @since 3.0.0 + * @since 2.3 */ public void setTimestampExpression(Expression timestampExpression) { this.timestampExpression = timestampExpression; } + /** + * Set the header mapper to use. + * @param headerMapper the mapper; can be null to disable header mapping. + * @since 2.3 + */ + public void setHeaderMapper(KafkaHeaderMapper headerMapper) { + this.headerMapper = headerMapper; + } + public KafkaTemplate getKafkaTemplate() { return this.kafkaTemplate; } @@ -167,7 +188,14 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { payload = null; } - ListenableFuture future = this.kafkaTemplate.send(topic, partitionId, timestamp, (K) messageKey, payload); + Headers headers = null; + if (this.headerMapper != null) { + headers = new RecordHeaders(); + this.headerMapper.fromHeaders(message.getHeaders(), headers); + } + ProducerRecord producerRecord = new ProducerRecord(topic, partitionId, timestamp, (K) messageKey, + payload, headers); + ListenableFuture future = this.kafkaTemplate.send(producerRecord); if (this.sync) { Long sendTimeout = this.sendTimeoutExpression.getValue(this.evaluationContext, message, Long.class); diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageDrivenAdapterTests.java b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageDrivenAdapterTests.java index a4638cbd44..88cb6bfd4a 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageDrivenAdapterTests.java +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageDrivenAdapterTests.java @@ -20,6 +20,7 @@ import static org.assertj.core.api.Assertions.assertThat; import java.lang.reflect.Type; import java.util.Arrays; +import java.util.Collections; import java.util.List; import java.util.Map; @@ -27,6 +28,8 @@ 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.producer.ProducerRecord; +import org.apache.kafka.common.header.Headers; +import org.apache.kafka.common.header.internals.RecordHeaders; import org.junit.ClassRule; import org.junit.Test; @@ -43,6 +46,7 @@ import org.springframework.kafka.core.ProducerFactory; import org.springframework.kafka.listener.KafkaMessageListenerContainer; import org.springframework.kafka.listener.config.ContainerProperties; import org.springframework.kafka.support.Acknowledgment; +import org.springframework.kafka.support.DefaultKafkaHeaderMapper; import org.springframework.kafka.support.KafkaHeaders; import org.springframework.kafka.support.KafkaNull; import org.springframework.kafka.support.converter.BatchMessageConverter; @@ -366,7 +370,10 @@ public class MessageDrivenAdapterTests { ProducerFactory pf = new DefaultKafkaProducerFactory(senderProps); KafkaTemplate template = new KafkaTemplate<>(pf); template.setDefaultTopic(topic3); - template.sendDefault(0, 1487694048607L, 1, "{\"bar\":\"baz\"}"); + Headers kHeaders = new RecordHeaders(); + MessageHeaders siHeaders = new MessageHeaders(Collections.singletonMap("foo", "bar")); + new DefaultKafkaHeaderMapper().fromHeaders(siHeaders, kHeaders); + template.send(new ProducerRecord<>(topic3, 0, 1487694048607L, 1, "{\"bar\":\"baz\"}", kHeaders)); Message received = out.receive(10000); assertThat(received).isNotNull(); @@ -379,6 +386,7 @@ public class MessageDrivenAdapterTests { assertThat(headers.get(KafkaHeaders.RECEIVED_TIMESTAMP)).isEqualTo(1487694048607L); assertThat(headers.get(KafkaHeaders.TIMESTAMP_TYPE)).isEqualTo("CREATE_TIME"); + assertThat(headers.get("foo")).isEqualTo("bar"); assertThat(received.getPayload()).isInstanceOf(Map.class); adapter.setPayloadType(Foo.class); diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandlerTests.java b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandlerTests.java index dd81f55049..7f982f8e9c 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandlerTests.java +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/outbound/KafkaProducerMessageHandlerTests.java @@ -23,6 +23,9 @@ import static org.springframework.kafka.test.assertj.KafkaConditions.partition; import static org.springframework.kafka.test.assertj.KafkaConditions.timestamp; import static org.springframework.kafka.test.assertj.KafkaConditions.value; +import java.util.HashMap; +import java.util.Map; + import org.apache.kafka.clients.consumer.Consumer; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.junit.BeforeClass; @@ -39,6 +42,7 @@ import org.springframework.kafka.core.DefaultKafkaConsumerFactory; import org.springframework.kafka.core.DefaultKafkaProducerFactory; import org.springframework.kafka.core.KafkaTemplate; import org.springframework.kafka.core.ProducerFactory; +import org.springframework.kafka.support.DefaultKafkaHeaderMapper; import org.springframework.kafka.support.KafkaHeaders; import org.springframework.kafka.support.KafkaNull; import org.springframework.kafka.test.rule.KafkaEmbedded; @@ -138,6 +142,7 @@ public class KafkaProducerMessageHandlerTests { .setHeader(KafkaHeaders.MESSAGE_KEY, 2) .setHeader(KafkaHeaders.PARTITION_ID, 1) .setHeader(KafkaHeaders.TIMESTAMP, 1487694048607L) + .setHeader("baz", "qux") .build(); handler.handleMessage(message); @@ -146,6 +151,10 @@ public class KafkaProducerMessageHandlerTests { assertThat(record).has(partition(1)); assertThat(record).has(value("foo")); assertThat(record).has(timestamp(1487694048607L)); + Map headers = new HashMap<>(); + new DefaultKafkaHeaderMapper().toHeaders(record.headers(), headers); + assertThat(headers.size()).isEqualTo(1); + assertThat(headers.get("baz")).isEqualTo("qux"); } @Test