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 4cfb9e49cd..e861fe389e 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 @@ -42,7 +42,7 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { private volatile Expression messageKeyExpression; - private volatile Expression partitionExpression; + private volatile Expression partitionIdExpression; @SuppressWarnings("unchecked") public KafkaProducerMessageHandler(final KafkaProducerContext kafkaProducerContext) { @@ -57,8 +57,8 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { this.messageKeyExpression = messageKeyExpression; } - public void setPartitionExpression(Expression partitionExpression) { - this.partitionExpression = partitionExpression; + public void setPartitionIdExpression(Expression partitionIdExpression) { + this.partitionIdExpression = partitionIdExpression; } public KafkaProducerContext getKafkaProducerContext() { @@ -77,8 +77,8 @@ public class KafkaProducerMessageHandler extends AbstractMessageHandler { this.topicExpression.getValue(this.evaluationContext, message, String.class) : message.getHeaders().get(KafkaHeaders.TOPIC, String.class); - Integer partitionId = this.partitionExpression != null ? - this.partitionExpression.getValue(this.evaluationContext, message, Integer.class) + Integer partitionId = this.partitionIdExpression != null ? + this.partitionIdExpression.getValue(this.evaluationContext, message, Integer.class) : message.getHeaders().get(KafkaHeaders.PARTITION_ID, Integer.class); Object messageKey = this.messageKeyExpression != null diff --git a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/support/ProducerConfiguration.java b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/support/ProducerConfiguration.java index 3a8eca19d0..cb5baf86c5 100644 --- a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/support/ProducerConfiguration.java +++ b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/support/ProducerConfiguration.java @@ -80,7 +80,13 @@ public class ProducerConfiguration { public Future send(String topic, Integer partition, K messageKey, V messagePayload) { String targetTopic = StringUtils.hasText(topic) ? topic : this.producerMetadata.getTopic(); - Future future = this.producer.send(new ProducerRecord<>(targetTopic, partition, messageKey, messagePayload)); + //If partition is sent by producer context then that takes precedence over custom partitioner. + if (partition == null && this.getProducerMetadata().getPartitioner() != null) { + partition = this.getProducerMetadata().getPartitioner().partition(messageKey, + this.producer.partitionsFor(targetTopic).size()); + } + Future future = + this.producer.send(new ProducerRecord<>(targetTopic, partition, messageKey, messagePayload)); if (!producerMetadata.isSync()) { return future; diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests-context.xml b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests-context.xml index f054e028f7..b63970beb6 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests-context.xml +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests-context.xml @@ -19,7 +19,9 @@ channel="inputToKafka" order="3" topic="foo" - message-key-expression="'bar'"> + message-key-expression="'bar'" + partition-id-expression="2" + > diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests.java b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests.java index 7cd9cb14fa..1392501a21 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests.java +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/config/xml/KafkaOutboundAdapterParserTests.java @@ -59,6 +59,7 @@ public class KafkaOutboundAdapterParserTests { assertEquals(messageHandler.getOrder(), 3); assertEquals("foo", TestUtils.getPropertyValue(messageHandler, "topicExpression.literalValue")); assertEquals("'bar'", TestUtils.getPropertyValue(messageHandler, "messageKeyExpression.expression")); + assertEquals("2", TestUtils.getPropertyValue(messageHandler, "partitionIdExpression.expression")); KafkaProducerContext producerContext = messageHandler.getKafkaProducerContext(); assertNotNull(producerContext); assertEquals(producerContext.getProducerConfigurations().size(), 2); diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/support/ProducerConfigurationTests.java b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/support/ProducerConfigurationTests.java index 982a90ac3e..a58589c18c 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/support/ProducerConfigurationTests.java +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/support/ProducerConfigurationTests.java @@ -18,9 +18,12 @@ package org.springframework.integration.kafka.support; import java.io.ByteArrayInputStream; import java.io.ObjectInputStream; +import java.util.ArrayList; +import kafka.producer.Partitioner; import org.apache.kafka.clients.producer.Producer; import org.apache.kafka.clients.producer.ProducerRecord; +import org.apache.kafka.common.PartitionInfo; import org.apache.kafka.common.serialization.ByteArraySerializer; import org.apache.kafka.common.serialization.StringSerializer; import org.junit.Assert; @@ -39,7 +42,6 @@ import org.springframework.messaging.Message; import org.springframework.messaging.support.GenericMessage; import kafka.serializer.DefaultEncoder; -import kafka.serializer.StringEncoder; /** * @author Soby Chacko @@ -242,6 +244,65 @@ public class ProducerConfigurationTests { Assert.assertEquals(capturedKeyMessage.topic(), "test"); } + @Test + @SuppressWarnings("unchecked") + public void testSendMessageWithNonDefaultKeyAndValueEncodersSpecifyingPartition() throws Exception { + final ProducerMetadata producerMetadata = new ProducerMetadata("test", String.class, String.class, new StringSerializer(), new StringSerializer()); + final Producer producer = Mockito.mock(Producer.class); + + final ProducerConfiguration configuration = + new ProducerConfiguration(producerMetadata, producer); + + configuration.send("test", 1, "key", "test message"); + + Mockito.verify(producer, Mockito.times(1)).send(Mockito.any(ProducerRecord.class)); + + final ArgumentCaptor> argument = + (ArgumentCaptor>) (Object) + ArgumentCaptor.forClass(ProducerRecord.class); + Mockito.verify(producer).send(argument.capture()); + + final ProducerRecord capturedKeyMessage = argument.getValue(); + + Assert.assertEquals(capturedKeyMessage.key(), "key"); + Assert.assertEquals(capturedKeyMessage.value(), "test message"); + Assert.assertEquals(capturedKeyMessage.topic(), "test"); + Assert.assertEquals(capturedKeyMessage.partition().intValue(), 1); + } + + @Test + @SuppressWarnings("unchecked") + public void testSendMessageWithNonDefaultKeyAndValueEncodersCustomPartitioner() throws Exception { + Partitioner customPartitioner = Mockito.mock(Partitioner.class); + final ProducerMetadata producerMetadata = new ProducerMetadata("test", String.class, String.class, new StringSerializer(), new StringSerializer()); + producerMetadata.setPartitioner(customPartitioner); + final Producer producer = Mockito.mock(Producer.class); + + final ProducerConfiguration configuration = + new ProducerConfiguration(producerMetadata, producer); + + ArrayList partitionInfos = new ArrayList<>(); + partitionInfos.add(Mockito.mock(PartitionInfo.class)); + Mockito.when(producer.partitionsFor("test")).thenReturn(partitionInfos); + Mockito.when(customPartitioner.partition("key", 1)).thenReturn(4); + + configuration.send("test", null, "key", "test message"); + + Mockito.verify(producer, Mockito.times(1)).send(Mockito.any(ProducerRecord.class)); + + final ArgumentCaptor> argument = + (ArgumentCaptor>) (Object) + ArgumentCaptor.forClass(ProducerRecord.class); + Mockito.verify(producer).send(argument.capture()); + + final ProducerRecord capturedKeyMessage = argument.getValue(); + + Assert.assertEquals(capturedKeyMessage.key(), "key"); + Assert.assertEquals(capturedKeyMessage.value(), "test message"); + Assert.assertEquals(capturedKeyMessage.topic(), "test"); + Assert.assertEquals(capturedKeyMessage.partition().intValue(), 4); + } + /** * User does not set an explicit key/value encoder, but send non-serializable object for both key/value */