diff --git a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/KafkaStreamsFunctionProcessor.java b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/KafkaStreamsFunctionProcessor.java index cae881178..284569738 100644 --- a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/KafkaStreamsFunctionProcessor.java +++ b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/KafkaStreamsFunctionProcessor.java @@ -26,6 +26,7 @@ import java.util.List; import java.util.Map; import java.util.Set; import java.util.TreeSet; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -110,7 +111,8 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro int inputCount = 1; ResolvableType currentOutputGeneric; - if (resolvableType.getRawClass().isAssignableFrom(BiFunction.class)) { + if (resolvableType.getRawClass().isAssignableFrom(BiFunction.class) || + resolvableType.getRawClass().isAssignableFrom(BiConsumer.class)) { inputCount = 2; currentOutputGeneric = resolvableType.getGeneric(2); } @@ -130,7 +132,8 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro popuateResolvableTypeMap(resolvableType, resolvableTypeMap, iterator); ResolvableType iterableResType = resolvableType; - int i = resolvableType.getRawClass().isAssignableFrom(BiFunction.class) ? 2 : 1; + int i = resolvableType.getRawClass().isAssignableFrom(BiFunction.class) || + resolvableType.getRawClass().isAssignableFrom(BiConsumer.class) ? 2 : 1; if (i == inputCount) { outboundResolvableType = iterableResType.getGeneric(i); } @@ -157,7 +160,9 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro private void popuateResolvableTypeMap(ResolvableType resolvableType, Map resolvableTypeMap, Iterator iterator) { final String next = iterator.next(); resolvableTypeMap.put(next, resolvableType.getGeneric(0)); - if (resolvableType.getRawClass() != null && resolvableType.getRawClass().isAssignableFrom(BiFunction.class) + if (resolvableType.getRawClass() != null && + (resolvableType.getRawClass().isAssignableFrom(BiFunction.class) || + resolvableType.getRawClass().isAssignableFrom(BiConsumer.class)) && iterator.hasNext()) { resolvableTypeMap.put(iterator.next(), resolvableType.getGeneric(1)); } @@ -175,6 +180,12 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro "No corresponding consumer beans found in the catalog"); consumer.accept(adaptedInboundArguments[0]); } + else if (resolvableType.getRawClass() != null && resolvableType.getRawClass().equals(BiConsumer.class)) { + BiConsumer biConsumer = (BiConsumer) this.beanFactory.getBean(functionName); + Assert.isTrue(biConsumer != null, + "No corresponding biConsumer beans found"); + biConsumer.accept(adaptedInboundArguments[0], adaptedInboundArguments[1]); + } else { Object result; if (resolvableType.getRawClass() != null && resolvableType.getRawClass().equals(BiFunction.class)) { diff --git a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/FunctionDetectorCondition.java b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/FunctionDetectorCondition.java index 339a5f010..1ce21466e 100644 --- a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/FunctionDetectorCondition.java +++ b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/FunctionDetectorCondition.java @@ -19,6 +19,7 @@ package org.springframework.cloud.stream.binder.kafka.streams.function; import java.lang.reflect.Method; import java.util.HashMap; import java.util.Map; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -51,6 +52,7 @@ public class FunctionDetectorCondition extends SpringBootCondition { Map functionTypes = context.getBeanFactory().getBeansOfType(Function.class); functionTypes.putAll(context.getBeanFactory().getBeansOfType(Consumer.class)); functionTypes.putAll(context.getBeanFactory().getBeansOfType(BiFunction.class)); + functionTypes.putAll(context.getBeanFactory().getBeansOfType(BiConsumer.class)); final Map kstreamFunctions = pruneFunctionBeansForKafkaStreams(functionTypes, context); if (!kstreamFunctions.isEmpty()) { diff --git a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsBindableProxyFactory.java b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsBindableProxyFactory.java index d91cc4471..037b8397f 100644 --- a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsBindableProxyFactory.java +++ b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsBindableProxyFactory.java @@ -22,6 +22,7 @@ import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Set; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -106,7 +107,8 @@ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFacto bindInput(argument, next); if (this.type.getRawClass() != null && - this.type.getRawClass().isAssignableFrom(BiFunction.class)) { + (this.type.getRawClass().isAssignableFrom(BiFunction.class) || + this.type.getRawClass().isAssignableFrom(BiConsumer.class))) { argument = this.type.getGeneric(resolvableTypeDepthCounter++); next = iterator.next(); bindInput(argument, next); @@ -172,7 +174,8 @@ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFacto return inputs; } int numberOfInputs = this.type.getRawClass() != null && - this.type.getRawClass().isAssignableFrom(BiFunction.class) ? 2 : getNumberOfInputs(); + (this.type.getRawClass().isAssignableFrom(BiFunction.class) || + this.type.getRawClass().isAssignableFrom(BiConsumer.class)) ? 2 : getNumberOfInputs(); if (numberOfInputs == 1) { inputs.add(String.format("%s_%s", this.functionName, DEFAULT_INPUT_SUFFIX)); return inputs; diff --git a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsFunctionBeanPostProcessor.java b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsFunctionBeanPostProcessor.java index 3b097a9a7..eb10aa2fd 100644 --- a/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsFunctionBeanPostProcessor.java +++ b/spring-cloud-stream-binder-kafka-streams/src/main/java/org/springframework/cloud/stream/binder/kafka/streams/function/KafkaStreamsFunctionBeanPostProcessor.java @@ -19,6 +19,7 @@ package org.springframework.cloud.stream.binder.kafka.streams.function; import java.lang.reflect.Method; import java.util.Map; import java.util.TreeMap; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -53,9 +54,11 @@ public class KafkaStreamsFunctionBeanPostProcessor implements InitializingBean, String[] functionNames = this.beanFactory.getBeanNamesForType(Function.class); String[] biFunctionNames = this.beanFactory.getBeanNamesForType(BiFunction.class); String[] consumerNames = this.beanFactory.getBeanNamesForType(Consumer.class); + String[] biConsumerNames = this.beanFactory.getBeanNamesForType(BiConsumer.class); Stream.concat( - Stream.concat(Stream.of(functionNames), Stream.of(consumerNames)), Stream.of(biFunctionNames)) + Stream.concat(Stream.of(functionNames), Stream.of(consumerNames)), + Stream.concat(Stream.of(biFunctionNames), Stream.of(biConsumerNames))) .forEach(this::extractResolvableTypes); } diff --git a/spring-cloud-stream-binder-kafka-streams/src/test/java/org/springframework/cloud/stream/binder/kafka/streams/function/StreamToTableJoinFunctionTests.java b/spring-cloud-stream-binder-kafka-streams/src/test/java/org/springframework/cloud/stream/binder/kafka/streams/function/StreamToTableJoinFunctionTests.java index 4b71f8a74..59b718ce3 100644 --- a/spring-cloud-stream-binder-kafka-streams/src/test/java/org/springframework/cloud/stream/binder/kafka/streams/function/StreamToTableJoinFunctionTests.java +++ b/spring-cloud-stream-binder-kafka-streams/src/test/java/org/springframework/cloud/stream/binder/kafka/streams/function/StreamToTableJoinFunctionTests.java @@ -20,6 +20,9 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.List; import java.util.Map; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Function; @@ -54,6 +57,7 @@ import org.springframework.kafka.core.KafkaTemplate; import org.springframework.kafka.test.EmbeddedKafkaBroker; import org.springframework.kafka.test.rule.EmbeddedKafkaRule; import org.springframework.kafka.test.utils.KafkaTestUtils; +import org.springframework.util.Assert; import static org.assertj.core.api.Assertions.assertThat; @@ -101,6 +105,77 @@ public class StreamToTableJoinFunctionTests { runTest(app, consumer); } + @Test + public void testStreamToTableBiConsumer() throws Exception { + SpringApplication app = new SpringApplication(BiConsumerApplication.class); + app.setWebApplicationType(WebApplicationType.NONE); + + Consumer consumer; + Map consumerProps = KafkaTestUtils.consumerProps("group-2", + "false", embeddedKafka); + consumerProps.put(ConsumerConfig.AUTO_OFFSET_RESET_CONFIG, "earliest"); + consumerProps.put(ConsumerConfig.KEY_DESERIALIZER_CLASS_CONFIG, StringDeserializer.class); + consumerProps.put(ConsumerConfig.VALUE_DESERIALIZER_CLASS_CONFIG, LongDeserializer.class); + DefaultKafkaConsumerFactory cf = new DefaultKafkaConsumerFactory<>(consumerProps); + consumer = cf.createConsumer(); + embeddedKafka.consumeFromAnEmbeddedTopic(consumer, "output-topic-1"); + + try (ConfigurableApplicationContext ignored = app.run("--server.port=0", + "--spring.jmx.enabled=false", + "--spring.cloud.stream.bindings.process_in_0.destination=user-clicks-1", + "--spring.cloud.stream.bindings.process_in_1.destination=user-regions-1", + "--spring.cloud.stream.kafka.streams.binder.configuration.default.key.serde" + + "=org.apache.kafka.common.serialization.Serdes$StringSerde", + "--spring.cloud.stream.kafka.streams.binder.configuration.default.value.serde" + + "=org.apache.kafka.common.serialization.Serdes$StringSerde", + "--spring.cloud.stream.kafka.streams.binder.configuration.commit.interval.ms=10000", + "--spring.cloud.stream.kafka.streams.bindings.process_in_0.consumer.applicationId" + + "=testStreamToTableBiConsumer", + "--spring.cloud.stream.kafka.streams.binder.brokers=" + embeddedKafka.getBrokersAsString())) { + + // Input 1: Region per user (multiple records allowed per user). + List> userRegions = Arrays.asList( + new KeyValue<>("alice", "asia") + ); + + Map senderProps1 = KafkaTestUtils.producerProps(embeddedKafka); + senderProps1.put(ProducerConfig.KEY_SERIALIZER_CLASS_CONFIG, StringSerializer.class); + senderProps1.put(ProducerConfig.VALUE_SERIALIZER_CLASS_CONFIG, StringSerializer.class); + + DefaultKafkaProducerFactory pf1 = new DefaultKafkaProducerFactory<>(senderProps1); + KafkaTemplate template1 = new KafkaTemplate<>(pf1, true); + template1.setDefaultTopic("user-regions-1"); + + for (KeyValue keyValue : userRegions) { + template1.sendDefault(keyValue.key, keyValue.value); + } + + // Input 2: Clicks per user (multiple records allowed per user). + List> userClicks = Arrays.asList( + new KeyValue<>("alice", 13L) + ); + + Map senderProps = KafkaTestUtils.producerProps(embeddedKafka); + senderProps.put(ProducerConfig.KEY_SERIALIZER_CLASS_CONFIG, StringSerializer.class); + senderProps.put(ProducerConfig.VALUE_SERIALIZER_CLASS_CONFIG, LongSerializer.class); + + DefaultKafkaProducerFactory pf = new DefaultKafkaProducerFactory<>(senderProps); + KafkaTemplate template = new KafkaTemplate<>(pf, true); + template.setDefaultTopic("user-clicks-1"); + + for (KeyValue keyValue : userClicks) { + template.sendDefault(keyValue.key, keyValue.value); + } + + Assert.isTrue(BiConsumerApplication.latch.await(10, TimeUnit.SECONDS), "Failed to receive message"); + + } + finally { + consumer.close(); + } + } + + private void runTest(SpringApplication app, Consumer consumer) { try (ConfigurableApplicationContext ignored = app.run("--server.port=0", "--spring.jmx.enabled=false", @@ -397,4 +472,18 @@ public class StreamToTableJoinFunctionTests { } } + @EnableAutoConfiguration + public static class BiConsumerApplication { + + static CountDownLatch latch = new CountDownLatch(2); + + @Bean + public BiConsumer, KTable> process() { + return (userClicksStream, userRegionsTable) -> { + userClicksStream.foreach((key, value) -> latch.countDown()); + userRegionsTable.toStream().foreach((key, value) -> latch.countDown()); + }; + } + } + }