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 d9b1e6e5e..cbd6630a4 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.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -108,31 +109,49 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro } private Map buildTypeMap(ResolvableType resolvableType) { - int inputCount = 1; - - ResolvableType resolvableTypeGeneric = resolvableType.getGeneric(1); - while (resolvableTypeGeneric != null && resolvableTypeGeneric.getRawClass() != null && (functionOrConsumerFound(resolvableTypeGeneric))) { - inputCount++; - resolvableTypeGeneric = resolvableTypeGeneric.getGeneric(1); - } - - final Set inputs = new LinkedHashSet<>(origInputs); Map resolvableTypeMap = new LinkedHashMap<>(); - final Iterator iterator = inputs.iterator(); + if (resolvableType != null && resolvableType.getRawClass() != null) { + int inputCount = 1; - popuateResolvableTypeMap(resolvableType, resolvableTypeMap, iterator); + ResolvableType currentOutputGeneric; + if (resolvableType.getRawClass().isAssignableFrom(BiFunction.class)) { + inputCount = 2; + currentOutputGeneric = resolvableType.getGeneric(2); + } + else { + currentOutputGeneric = resolvableType.getGeneric(1); + } + while (currentOutputGeneric != null && currentOutputGeneric.getRawClass() != null + && (functionOrConsumerFound(currentOutputGeneric))) { + inputCount++; + currentOutputGeneric = currentOutputGeneric.getGeneric(1); + } - ResolvableType iterableResType = resolvableType; - for (int i = 1; i < inputCount; i++) { - if (iterator.hasNext()) { - iterableResType = iterableResType.getGeneric(1); - if (iterableResType.getRawClass() != null && - functionOrConsumerFound(iterableResType)) { - popuateResolvableTypeMap(iterableResType, resolvableTypeMap, iterator); + final Set inputs = new LinkedHashSet<>(origInputs); + + final Iterator iterator = inputs.iterator(); + + popuateResolvableTypeMap(resolvableType, resolvableTypeMap, iterator); + + ResolvableType iterableResType = resolvableType; + int i = resolvableType.getRawClass().isAssignableFrom(BiFunction.class) ? 2 : 1; + if (i == inputCount) { + outboundResolvableType = iterableResType.getGeneric(i); + } + else { + while (i < inputCount) { + if (iterator.hasNext()) { + iterableResType = iterableResType.getGeneric(1); + if (iterableResType.getRawClass() != null && + functionOrConsumerFound(iterableResType)) { + popuateResolvableTypeMap(iterableResType, resolvableTypeMap, iterator); + } + i++; + } } + outboundResolvableType = iterableResType.getGeneric(1); } } - outboundResolvableType = iterableResType.getGeneric(1); return resolvableTypeMap; } @@ -144,6 +163,10 @@ 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) + && iterator.hasNext()) { + resolvableTypeMap.put(iterator.next(), resolvableType.getGeneric(1)); + } origInputs.remove(next); } @@ -159,9 +182,17 @@ public class KafkaStreamsFunctionProcessor extends AbstractKafkaStreamsBinderPro consumer.accept(adaptedInboundArguments[0]); } else { - Function function = (Function) beanFactory.getBean(functionName); - Assert.isTrue(function != null, "Function bean cannot be null"); - Object result = function.apply(adaptedInboundArguments[0]); + Object result; + if (resolvableType.getRawClass() != null && resolvableType.getRawClass().equals(BiFunction.class)) { + BiFunction biFunction = (BiFunction) beanFactory.getBean(functionName); + Assert.isTrue(biFunction != null, "Biunction bean cannot be null"); + result = biFunction.apply(adaptedInboundArguments[0], adaptedInboundArguments[1]); + } + else { + Function function = (Function) beanFactory.getBean(functionName); + Assert.isTrue(function != null, "Function bean cannot be null"); + result = function.apply(adaptedInboundArguments[0]); + } int i = 1; while (result instanceof Function || result instanceof Consumer) { if (result instanceof Function) { 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 18f0fd535..339a5f010 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.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -49,16 +50,17 @@ public class FunctionDetectorCondition extends SpringBootCondition { if (context != null && context.getBeanFactory() != null) { Map functionTypes = context.getBeanFactory().getBeansOfType(Function.class); functionTypes.putAll(context.getBeanFactory().getBeansOfType(Consumer.class)); + functionTypes.putAll(context.getBeanFactory().getBeansOfType(BiFunction.class)); final Map kstreamFunctions = pruneFunctionBeansForKafkaStreams(functionTypes, context); if (!kstreamFunctions.isEmpty()) { - return ConditionOutcome.match("Matched. Function/Consumer beans found"); + return ConditionOutcome.match("Matched. Function/BiFunction/Consumer beans found"); } else { - return ConditionOutcome.noMatch("No match. No Function/Consumer beans found"); + return ConditionOutcome.noMatch("No match. No Function/BiFunction/Consumer beans found"); } } - return ConditionOutcome.noMatch("No match. No Function/Consumer beans found"); + return ConditionOutcome.noMatch("No match. No Function/BiFunction/Consumer beans found"); } private static Map pruneFunctionBeansForKafkaStreams(Map originalFunctionBeans, 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 48516b354..461881d42 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.BiFunction; import java.util.function.Consumer; import java.util.function.Function; @@ -67,6 +68,8 @@ import org.springframework.util.CollectionUtils; */ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFactory implements InitializingBean, BeanFactoryAware { + private static final String DEFAULT_INPUT_SUFFIX = "input"; + private static Log log = LogFactory.getLog(BindableProxyFactory.class); @Autowired @@ -90,57 +93,58 @@ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFacto Assert.notEmpty(KafkaStreamsBindableProxyFactory.this.bindingTargetFactories, "'bindingTargetFactories' cannot be empty"); - ResolvableType arg0 = this.type.getGeneric(0); + int resolvableTypeDepthCounter = 0; + ResolvableType argument = this.type.getGeneric(resolvableTypeDepthCounter++); List inputBindings = buildInputBindings(); Iterator iterator = inputBindings.iterator(); String next = iterator.next(); - bindInput(arg0, next); - BeanDefinitionRegistry registry = (BeanDefinitionRegistry) beanFactory; + bindInput(argument, next); - RootBeanDefinition rootBeanDefinition = new RootBeanDefinition(); - rootBeanDefinition.setInstanceSupplier(() -> inputHolders.get(next).getBoundTarget()); - registry.registerBeanDefinition(next, rootBeanDefinition); + if (this.type.getRawClass() != null && + this.type.getRawClass().isAssignableFrom(BiFunction.class)) { + argument = this.type.getGeneric(resolvableTypeDepthCounter++); + next = iterator.next(); + bindInput(argument, next); + } + ResolvableType outboundArgument = this.type.getGeneric(resolvableTypeDepthCounter); - ResolvableType arg1 = this.type.getGeneric(1); - - while (isAnotherFunctionOrConsumerFound(arg1)) { - arg0 = arg1.getGeneric(0); + while (isAnotherFunctionOrConsumerFound(outboundArgument)) { + //The function is a curried function. We should introspect the partial function chain hierarchy. + argument = outboundArgument.getGeneric(0); String next1 = iterator.next(); - bindInput(arg0, next1); - RootBeanDefinition rootBeanDefinition1 = new RootBeanDefinition(); - rootBeanDefinition1.setInstanceSupplier(() -> inputHolders.get(next1).getBoundTarget()); - registry.registerBeanDefinition(next1, rootBeanDefinition1); - - arg1 = arg1.getGeneric(1); + bindInput(argument, next1); + outboundArgument = outboundArgument.getGeneric(1); } //Introspect output for binding. - if (arg1 != null && arg1.getRawClass() != null && (arg1.isArray() || arg1.getRawClass().isAssignableFrom(KStream.class))) { + if (outboundArgument != null && outboundArgument.getRawClass() != null && (!outboundArgument.isArray() && + outboundArgument.getRawClass().isAssignableFrom(KStream.class))) { // if the type is array, we need to do a late binding as we don't know the number of // output bindings at this point in the flow. - if (!arg1.isArray()) { - List outputBindings = streamFunctionProperties.getOutputBindings().get(this.functionName); - String outputBinding = null; - if (!CollectionUtils.isEmpty(outputBindings)) { - Iterator outputBindingsIter = outputBindings.iterator(); - if (outputBindingsIter.hasNext()) { - outputBinding = outputBindingsIter.next(); - } + List outputBindings = streamFunctionProperties.getOutputBindings().get(this.functionName); + String outputBinding = null; + if (!CollectionUtils.isEmpty(outputBindings)) { + Iterator outputBindingsIter = outputBindings.iterator(); + if (outputBindingsIter.hasNext()) { + outputBinding = outputBindingsIter.next(); } - else { - outputBinding = this.functionName + "-" + "output"; - } - Assert.isTrue(outputBinding != null, "output binding is not inferred."); - KafkaStreamsBindableProxyFactory.this.outputHolders.put(outputBinding, - new BoundTargetHolder(getBindingTargetFactory(KStream.class) - .createOutput(outputBinding), true)); - String outputBinding1 = outputBinding; - RootBeanDefinition rootBeanDefinition1 = new RootBeanDefinition(); - rootBeanDefinition1.setInstanceSupplier(() -> outputHolders.get(outputBinding1).getBoundTarget()); - registry.registerBeanDefinition(outputBinding1, rootBeanDefinition1); + } + else { + outputBinding = this.functionName + "-" + "output"; + } + Assert.isTrue(outputBinding != null, "output binding is not inferred."); + KafkaStreamsBindableProxyFactory.this.outputHolders.put(outputBinding, + new BoundTargetHolder(getBindingTargetFactory(KStream.class) + .createOutput(outputBinding), true)); + String outputBinding1 = outputBinding; + RootBeanDefinition rootBeanDefinition1 = new RootBeanDefinition(); + rootBeanDefinition1.setInstanceSupplier(() -> outputHolders.get(outputBinding1).getBoundTarget()); + BeanDefinitionRegistry registry = (BeanDefinitionRegistry) beanFactory; + registry.registerBeanDefinition(outputBinding1, rootBeanDefinition1); + } } @@ -163,15 +167,16 @@ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFacto inputs.addAll(inputBindings); return inputs; } - int numberOfInputs = getNumberOfInputs(); + int numberOfInputs = this.type.getRawClass() != null && + this.type.getRawClass().isAssignableFrom(BiFunction.class) ? 2 : getNumberOfInputs(); if (numberOfInputs == 1) { - inputs.add(this.functionName + "-" + "input"); + inputs.add(this.functionName + "-" + DEFAULT_INPUT_SUFFIX); return inputs; } else { int i = 0; while (i < numberOfInputs) { - inputs.add(this.functionName + "-" + "input" + "-" + i++); + inputs.add(this.functionName + "-" + DEFAULT_INPUT_SUFFIX + "-" + i++); } return inputs; } @@ -207,6 +212,13 @@ public class KafkaStreamsBindableProxyFactory extends AbstractBindableProxyFacto .createInput(inputName), true)); } } + + BeanDefinitionRegistry registry = (BeanDefinitionRegistry) beanFactory; + + RootBeanDefinition rootBeanDefinition = new RootBeanDefinition(); + rootBeanDefinition.setInstanceSupplier(() -> inputHolders.get(inputName).getBoundTarget()); + registry.registerBeanDefinition(inputName, rootBeanDefinition); + } @Override 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 9730d47ab..fd8fc1a74 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.BiFunction; import java.util.function.Consumer; import java.util.function.Function; import java.util.stream.Stream; @@ -58,9 +59,12 @@ public class KafkaStreamsFunctionBeanPostProcessor implements InitializingBean, public void afterPropertiesSet() { String[] functionNames = this.beanFactory.getBeanNamesForType(Function.class); + String[] biFunctionNames = this.beanFactory.getBeanNamesForType(BiFunction.class); String[] consumerNames = this.beanFactory.getBeanNamesForType(Consumer.class); - Stream.concat(Stream.of(functionNames), Stream.of(consumerNames)).forEach(this::extractResolvableTypes); + Stream.concat( + Stream.concat(Stream.of(functionNames), Stream.of(consumerNames)), Stream.of(biFunctionNames)) + .forEach(this::extractResolvableTypes); BindableProvider bindableProvider = clazz -> clazz.isAssignableFrom(KStream.class) || clazz.isAssignableFrom(KTable.class) 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 01381aa47..bf7328092 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,7 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.List; import java.util.Map; +import java.util.function.BiFunction; import java.util.function.Function; import org.apache.kafka.clients.consumer.Consumer; @@ -65,7 +66,7 @@ public class StreamToTableJoinFunctionTests { private static EmbeddedKafkaBroker embeddedKafka = embeddedKafkaRule.getEmbeddedKafka(); @Test - public void testStreamToTable() throws Exception { + public void testStreamToTable() { SpringApplication app = new SpringApplication(CountClicksPerRegionApplication.class); app.setWebApplicationType(WebApplicationType.NONE); @@ -79,6 +80,28 @@ public class StreamToTableJoinFunctionTests { consumer = cf.createConsumer(); embeddedKafka.consumeFromAnEmbeddedTopic(consumer, "output-topic-1"); + runTest(app, consumer); + } + + @Test + public void testStreamToTableBiFunction() { + SpringApplication app = new SpringApplication(BiFunctionCountClicksPerRegionApplication.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"); + + runTest(app, consumer); + } + + private void runTest(SpringApplication app, Consumer consumer) { try (ConfigurableApplicationContext ignored = app.run("--server.port=0", "--spring.jmx.enabled=false", "--spring.cloud.stream.bindings.process-input-0.destination=user-clicks-1", @@ -168,11 +191,11 @@ public class StreamToTableJoinFunctionTests { @Test public void testGlobalStartOffsetWithLatestAndIndividualBindingWthEarliest() throws Exception { - SpringApplication app = new SpringApplication(CountClicksPerRegionApplication.class); + SpringApplication app = new SpringApplication(BiFunctionCountClicksPerRegionApplication.class); app.setWebApplicationType(WebApplicationType.NONE); Consumer consumer; - Map consumerProps = KafkaTestUtils.consumerProps("group-2", + Map consumerProps = KafkaTestUtils.consumerProps("group-3", "false", embeddedKafka); consumerProps.put(ConsumerConfig.AUTO_OFFSET_RESET_CONFIG, "earliest"); consumerProps.put(ConsumerConfig.KEY_DESERIALIZER_CLASS_CONFIG, StringDeserializer.class); @@ -356,4 +379,22 @@ public class StreamToTableJoinFunctionTests { } } + @EnableAutoConfiguration + @EnableConfigurationProperties(KafkaStreamsApplicationSupportProperties.class) + public static class BiFunctionCountClicksPerRegionApplication { + + @Bean + public BiFunction, KTable, KStream> process() { + return (userClicksStream, userRegionsTable) -> (userClicksStream + .leftJoin(userRegionsTable, (clicks, region) -> new RegionWithClicks(region == null ? + "UNKNOWN" : region, clicks), + Joined.with(Serdes.String(), Serdes.Long(), null)) + .map((user, regionWithClicks) -> new KeyValue<>(regionWithClicks.getRegion(), + regionWithClicks.getClicks())) + .groupByKey(Serialized.with(Serdes.String(), Serdes.Long())) + .reduce(Long::sum) + .toStream()); + } + } + }