diff --git a/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/condition/OnBeanCondition.java b/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/condition/OnBeanCondition.java index b90ac437e9..521bbc7fb3 100644 --- a/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/condition/OnBeanCondition.java +++ b/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/condition/OnBeanCondition.java @@ -28,7 +28,9 @@ import java.util.LinkedHashSet; import java.util.List; import java.util.Locale; import java.util.Map; +import java.util.Map.Entry; import java.util.Set; +import java.util.function.Predicate; import org.springframework.aop.scope.ScopedProxyUtils; import org.springframework.beans.factory.BeanFactory; @@ -113,61 +115,90 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat @Override public ConditionOutcome getMatchOutcome(ConditionContext context, AnnotatedTypeMetadata metadata) { - ConditionMessage matchMessage = ConditionMessage.empty(); + ConditionOutcome matchOutcome = ConditionOutcome.match(); MergedAnnotations annotations = metadata.getAnnotations(); if (annotations.isPresent(ConditionalOnBean.class)) { Spec spec = new Spec<>(context, metadata, annotations, ConditionalOnBean.class); - MatchResult matchResult = getMatchingBeans(context, spec); - if (!matchResult.isAllMatched()) { - String reason = createOnBeanNoMatchReason(matchResult); - return ConditionOutcome.noMatch(spec.message().because(reason)); + matchOutcome = evaluateConditionalOnBean(spec, matchOutcome.getConditionMessage()); + if (!matchOutcome.isMatch()) { + return matchOutcome; } - matchMessage = spec.message(matchMessage) - .found("bean", "beans") - .items(Style.QUOTE, matchResult.getNamesOfAllMatches()); } if (metadata.isAnnotated(ConditionalOnSingleCandidate.class.getName())) { - Spec spec = new SingleCandidateSpec(context, metadata, annotations); - MatchResult matchResult = getMatchingBeans(context, spec); - if (!matchResult.isAllMatched()) { - return ConditionOutcome.noMatch(spec.message().didNotFind("any beans").atAll()); - } - Set allBeans = matchResult.getNamesOfAllMatches(); - if (allBeans.size() == 1) { - matchMessage = spec.message(matchMessage).found("a single bean").items(Style.QUOTE, allBeans); - } - else { - List primaryBeans = getPrimaryBeans(context.getBeanFactory(), allBeans, - spec.getStrategy() == SearchStrategy.ALL); - if (primaryBeans.isEmpty()) { - return ConditionOutcome - .noMatch(spec.message().didNotFind("a primary bean from beans").items(Style.QUOTE, allBeans)); - } - if (primaryBeans.size() > 1) { - return ConditionOutcome - .noMatch(spec.message().found("multiple primary beans").items(Style.QUOTE, primaryBeans)); - } - matchMessage = spec.message(matchMessage) - .found("a single primary bean '" + primaryBeans.get(0) + "' from beans") - .items(Style.QUOTE, allBeans); + Spec spec = new SingleCandidateSpec(context, metadata, + metadata.getAnnotations()); + matchOutcome = evaluateConditionalOnSingleCandidate(spec, matchOutcome.getConditionMessage()); + if (!matchOutcome.isMatch()) { + return matchOutcome; } } if (metadata.isAnnotated(ConditionalOnMissingBean.class.getName())) { Spec spec = new Spec<>(context, metadata, annotations, ConditionalOnMissingBean.class); - MatchResult matchResult = getMatchingBeans(context, spec); - if (matchResult.isAnyMatched()) { - String reason = createOnMissingBeanNoMatchReason(matchResult); - return ConditionOutcome.noMatch(spec.message().because(reason)); + matchOutcome = evaluateConditionalOnMissingBean(spec, matchOutcome.getConditionMessage()); + if (!matchOutcome.isMatch()) { + return matchOutcome; } - matchMessage = spec.message(matchMessage).didNotFind("any beans").atAll(); } - return ConditionOutcome.match(matchMessage); + return matchOutcome; } - protected final MatchResult getMatchingBeans(ConditionContext context, Spec spec) { - ClassLoader classLoader = context.getClassLoader(); - ConfigurableListableBeanFactory beanFactory = context.getBeanFactory(); + private ConditionOutcome evaluateConditionalOnBean(Spec spec, ConditionMessage matchMessage) { + MatchResult matchResult = getMatchingBeans(spec); + if (!matchResult.isAllMatched()) { + String reason = createOnBeanNoMatchReason(matchResult); + return ConditionOutcome.noMatch(spec.message().because(reason)); + } + return ConditionOutcome.match(spec.message(matchMessage) + .found("bean", "beans") + .items(Style.QUOTE, matchResult.getNamesOfAllMatches())); + } + + private ConditionOutcome evaluateConditionalOnSingleCandidate(Spec spec, + ConditionMessage matchMessage) { + MatchResult matchResult = getMatchingBeans(spec); + if (!matchResult.isAllMatched()) { + return ConditionOutcome.noMatch(spec.message().didNotFind("any beans").atAll()); + } + Set allBeans = matchResult.getNamesOfAllMatches(); + if (allBeans.size() == 1) { + return ConditionOutcome + .match(spec.message(matchMessage).found("a single bean").items(Style.QUOTE, allBeans)); + } + Map beanDefinitions = getBeanDefinitions(spec.context.getBeanFactory(), allBeans, + spec.getStrategy() == SearchStrategy.ALL); + List primaryBeans = getPrimaryBeans(beanDefinitions); + if (primaryBeans.size() == 1) { + return ConditionOutcome.match(spec.message(matchMessage) + .found("a single primary bean '" + primaryBeans.get(0) + "' from beans") + .items(Style.QUOTE, allBeans)); + } + if (primaryBeans.size() > 1) { + return ConditionOutcome + .noMatch(spec.message().found("multiple primary beans").items(Style.QUOTE, primaryBeans)); + } + List nonFallbackBeans = getNonFallbackBeans(beanDefinitions); + if (nonFallbackBeans.size() == 1) { + return ConditionOutcome.match(spec.message(matchMessage) + .found("a single non-fallback bean '" + nonFallbackBeans.get(0) + "' from beans") + .items(Style.QUOTE, allBeans)); + } + return ConditionOutcome.noMatch(spec.message().found("multiple beans").items(Style.QUOTE, allBeans)); + } + + private ConditionOutcome evaluateConditionalOnMissingBean(Spec spec, + ConditionMessage matchMessage) { + MatchResult matchResult = getMatchingBeans(spec); + if (matchResult.isAnyMatched()) { + String reason = createOnMissingBeanNoMatchReason(matchResult); + return ConditionOutcome.noMatch(spec.message().because(reason)); + } + return ConditionOutcome.match(spec.message(matchMessage).didNotFind("any beans").atAll()); + } + + protected final MatchResult getMatchingBeans(Spec spec) { + ClassLoader classLoader = spec.getContext().getClassLoader(); + ConfigurableListableBeanFactory beanFactory = spec.getContext().getBeanFactory(); boolean considerHierarchy = spec.getStrategy() != SearchStrategy.CURRENT; Set> parameterizedContainers = spec.getParameterizedContainers(); if (spec.getStrategy() == SearchStrategy.ANCESTORS) { @@ -373,16 +404,32 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat } } - private List getPrimaryBeans(ConfigurableListableBeanFactory beanFactory, Set beanNames, - boolean considerHierarchy) { - List primaryBeans = new ArrayList<>(); + private Map getBeanDefinitions(ConfigurableListableBeanFactory beanFactory, + Set beanNames, boolean considerHierarchy) { + Map definitions = new HashMap<>(beanNames.size()); for (String beanName : beanNames) { BeanDefinition beanDefinition = findBeanDefinition(beanFactory, beanName, considerHierarchy); - if (beanDefinition != null && beanDefinition.isPrimary()) { - primaryBeans.add(beanName); + definitions.put(beanName, beanDefinition); + } + return definitions; + } + + private List getPrimaryBeans(Map beanDefinitions) { + return getMatchingBeans(beanDefinitions, BeanDefinition::isPrimary); + } + + private List getNonFallbackBeans(Map beanDefinitions) { + return getMatchingBeans(beanDefinitions, Predicate.not(BeanDefinition::isFallback)); + } + + private List getMatchingBeans(Map beanDefinitions, Predicate test) { + List matches = new ArrayList<>(); + for (Entry namedBeanDefinition : beanDefinitions.entrySet()) { + if (test.test(namedBeanDefinition.getValue())) { + matches.add(namedBeanDefinition.getKey()); } } - return primaryBeans; + return matches; } private BeanDefinition findBeanDefinition(ConfigurableListableBeanFactory beanFactory, String beanName, @@ -420,7 +467,7 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat */ private static class Spec { - private final ClassLoader classLoader; + private final ConditionContext context; private final Class annotationType; @@ -442,7 +489,7 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat .filter(MergedAnnotationPredicates.unique(MergedAnnotation::getMetaTypes)) .collect(MergedAnnotationCollectors.toMultiValueMap(Adapt.CLASS_TO_STRING)); MergedAnnotation annotation = annotations.get(annotationType); - this.classLoader = context.getClassLoader(); + this.context = context; this.annotationType = annotationType; this.names = extract(attributes, "name"); this.annotations = extract(attributes, "annotation"); @@ -497,7 +544,7 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat Set> resolved = new LinkedHashSet<>(classNames.size()); for (String className : classNames) { try { - resolved.add(resolve(className, this.classLoader)); + resolved.add(resolve(className, this.context.getClassLoader())); } catch (ClassNotFoundException | NoClassDefFoundError ex) { // Ignore @@ -596,31 +643,35 @@ class OnBeanCondition extends FilteringSpringBootCondition implements Configurat return (this.strategy != null) ? this.strategy : SearchStrategy.ALL; } - Set getNames() { + private ConditionContext getContext() { + return this.context; + } + + private Set getNames() { return this.names; } - Set getTypes() { + protected Set getTypes() { return this.types; } - Set getAnnotations() { + private Set getAnnotations() { return this.annotations; } - Set getIgnoredTypes() { + private Set getIgnoredTypes() { return this.ignoredTypes; } - Set> getParameterizedContainers() { + private Set> getParameterizedContainers() { return this.parameterizedContainers; } - ConditionMessage.Builder message() { + private ConditionMessage.Builder message() { return ConditionMessage.forCondition(this.annotationType, this); } - ConditionMessage.Builder message(ConditionMessage message) { + private ConditionMessage.Builder message(ConditionMessage message) { return message.andCondition(this.annotationType, this); } diff --git a/spring-boot-project/spring-boot-autoconfigure/src/test/java/org/springframework/boot/autoconfigure/condition/ConditionalOnSingleCandidateTests.java b/spring-boot-project/spring-boot-autoconfigure/src/test/java/org/springframework/boot/autoconfigure/condition/ConditionalOnSingleCandidateTests.java index 292c119f11..7431fa2a33 100644 --- a/spring-boot-project/spring-boot-autoconfigure/src/test/java/org/springframework/boot/autoconfigure/condition/ConditionalOnSingleCandidateTests.java +++ b/spring-boot-project/spring-boot-autoconfigure/src/test/java/org/springframework/boot/autoconfigure/condition/ConditionalOnSingleCandidateTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2012-2023 the original author or authors. + * Copyright 2012-2024 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -21,6 +21,7 @@ import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.Fallback; import org.springframework.context.annotation.Primary; import org.springframework.context.annotation.Scope; import org.springframework.context.annotation.ScopedProxyMode; @@ -114,6 +115,17 @@ class ConditionalOnSingleCandidateTests { }); } + @Test + void singleCandidateTwoCandidatesOneNormalOneFallback() { + this.contextRunner + .withUserConfiguration(AlphaFallbackConfiguration.class, BravoConfiguration.class, + OnBeanSingleCandidateConfiguration.class) + .run((context) -> { + assertThat(context).hasBean("consumer"); + assertThat(context.getBean("consumer")).isEqualTo("bravo"); + }); + } + @Test void singleCandidateMultipleCandidatesMultiplePrimary() { this.contextRunner @@ -122,6 +134,14 @@ class ConditionalOnSingleCandidateTests { .run((context) -> assertThat(context).doesNotHaveBean("consumer")); } + @Test + void singleCandidateMultipleCandidatesAllFallback() { + this.contextRunner + .withUserConfiguration(AlphaFallbackConfiguration.class, BravoFallbackConfiguration.class, + OnBeanSingleCandidateConfiguration.class) + .run((context) -> assertThat(context).doesNotHaveBean("consumer")); + } + @Test void invalidAnnotationTwoTypes() { this.contextRunner.withUserConfiguration(OnBeanSingleCandidateTwoTypesConfiguration.class).run((context) -> { @@ -208,6 +228,17 @@ class ConditionalOnSingleCandidateTests { } + @Configuration(proxyBeanMethods = false) + static class AlphaFallbackConfiguration { + + @Bean + @Fallback + String alpha() { + return "alpha"; + } + + } + @Configuration(proxyBeanMethods = false) static class AlphaScopedProxyConfiguration { @@ -240,4 +271,15 @@ class ConditionalOnSingleCandidateTests { } + @Configuration(proxyBeanMethods = false) + static class BravoFallbackConfiguration { + + @Bean + @Fallback + String bravo() { + return "bravo"; + } + + } + }