diff --git a/spring-boot-project/spring-boot-test/src/main/java/org/springframework/boot/test/context/ImportsContextCustomizer.java b/spring-boot-project/spring-boot-test/src/main/java/org/springframework/boot/test/context/ImportsContextCustomizer.java index fc4a8babbf..ca5786240b 100644 --- a/spring-boot-project/spring-boot-test/src/main/java/org/springframework/boot/test/context/ImportsContextCustomizer.java +++ b/spring-boot-project/spring-boot-test/src/main/java/org/springframework/boot/test/context/ImportsContextCustomizer.java @@ -16,13 +16,11 @@ package org.springframework.boot.test.context; -import java.lang.annotation.Annotation; -import java.lang.reflect.AnnotatedElement; import java.lang.reflect.Constructor; import java.util.Collections; -import java.util.HashSet; import java.util.LinkedHashSet; import java.util.Set; +import java.util.stream.Collectors; import org.springframework.beans.BeansException; import org.springframework.beans.factory.BeanFactory; @@ -42,10 +40,9 @@ import org.springframework.context.annotation.ImportBeanDefinitionRegistrar; import org.springframework.context.annotation.ImportSelector; import org.springframework.context.support.AbstractApplicationContext; import org.springframework.core.Ordered; -import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.core.annotation.AnnotationFilter; import org.springframework.core.annotation.MergedAnnotation; import org.springframework.core.annotation.MergedAnnotations; -import org.springframework.core.annotation.MergedAnnotations.SearchStrategy; import org.springframework.core.annotation.Order; import org.springframework.core.style.ToStringCreator; import org.springframework.core.type.AnnotationMetadata; @@ -59,6 +56,7 @@ import org.springframework.util.ReflectionUtils; * * @author Phillip Webb * @author Andy Wilkinson + * @author Laurent Martelli * @see ImportsContextCustomizerFactory */ class ImportsContextCustomizer implements ContextCustomizer { @@ -217,68 +215,36 @@ class ImportsContextCustomizer implements ContextCustomizer { */ static class ContextCustomizerKey { - private static final Class[] NO_IMPORTS = {}; - private static final Set ANNOTATION_FILTERS; - static { - Set filters = new HashSet<>(); - filters.add(new JavaLangAnnotationFilter()); - filters.add(new KotlinAnnotationFilter()); - filters.add(new SpockAnnotationFilter()); - filters.add(new JUnitAnnotationFilter()); - ANNOTATION_FILTERS = Collections.unmodifiableSet(filters); + Set annotationFilters = new LinkedHashSet<>(); + annotationFilters.add(AnnotationFilter.PLAIN); + annotationFilters.add("kotlin.Metadata"::equals); + annotationFilters.add(AnnotationFilter.packages("kotlin.annotation")); + annotationFilters.add(AnnotationFilter.packages("org.spockframework", "spock")); + annotationFilters.add(AnnotationFilter.packages("org.junit")); + ANNOTATION_FILTERS = Collections.unmodifiableSet(annotationFilters); } - private final Set key; ContextCustomizerKey(Class testClass) { - Set annotations = new HashSet<>(); - Set> seen = new HashSet<>(); - collectClassAnnotations(testClass, annotations, seen); + MergedAnnotations annotations = MergedAnnotations.search(MergedAnnotations.SearchStrategy.TYPE_HIERARCHY) + .withAnnotationFilter(this::isFilteredAnnotation) + .from(testClass); Set determinedImports = determineImports(annotations, testClass); - this.key = Collections.unmodifiableSet((determinedImports != null) ? determinedImports : annotations); + this.key = (determinedImports != null) ? determinedImports : synthesize(annotations); } - private void collectClassAnnotations(Class classType, Set annotations, Set> seen) { - if (seen.add(classType)) { - collectElementAnnotations(classType, annotations, seen); - for (Class interfaceType : classType.getInterfaces()) { - collectClassAnnotations(interfaceType, annotations, seen); - } - if (classType.getSuperclass() != null) { - collectClassAnnotations(classType.getSuperclass(), annotations, seen); - } - } + private boolean isFilteredAnnotation(String typeName) { + return ANNOTATION_FILTERS.stream().anyMatch((filter) -> filter.matches(typeName)); } - private void collectElementAnnotations(AnnotatedElement element, Set annotations, - Set> seen) { - for (MergedAnnotation mergedAnnotation : MergedAnnotations.from(element, - SearchStrategy.DIRECT)) { - Annotation annotation = mergedAnnotation.synthesize(); - if (!isIgnoredAnnotation(annotation)) { - annotations.add(annotation); - collectClassAnnotations(annotation.annotationType(), annotations, seen); - } - } - } - - private boolean isIgnoredAnnotation(Annotation annotation) { - for (AnnotationFilter annotationFilter : ANNOTATION_FILTERS) { - if (annotationFilter.isIgnored(annotation)) { - return true; - } - } - return false; - } - - private Set determineImports(Set annotations, Class testClass) { + private Set determineImports(MergedAnnotations annotations, Class testClass) { Set determinedImports = new LinkedHashSet<>(); - AnnotationMetadata testClassMetadata = AnnotationMetadata.introspect(testClass); - for (Annotation annotation : annotations) { - for (Class source : getImports(annotation)) { - Set determinedSourceImports = determineImports(source, testClassMetadata); + AnnotationMetadata metadata = AnnotationMetadata.introspect(testClass); + for (MergedAnnotation annotation : annotations.stream(Import.class).toList()) { + for (Class source : annotation.getClassArray(MergedAnnotation.VALUE)) { + Set determinedSourceImports = determineImports(source, metadata); if (determinedSourceImports == null) { return null; } @@ -288,13 +254,6 @@ class ImportsContextCustomizer implements ContextCustomizer { return determinedImports; } - private Class[] getImports(Annotation annotation) { - if (annotation instanceof Import importAnnotation) { - return importAnnotation.value(); - } - return NO_IMPORTS; - } - private Set determineImports(Class source, AnnotationMetadata metadata) { if (DeterminableImports.class.isAssignableFrom(source)) { // We can determine the imports @@ -310,6 +269,10 @@ class ImportsContextCustomizer implements ContextCustomizer { return Collections.singleton(source.getName()); } + private Set synthesize(MergedAnnotations annotations) { + return annotations.stream().map(MergedAnnotation::synthesize).collect(Collectors.toSet()); + } + @SuppressWarnings("unchecked") private T instantiate(Class source) { try { @@ -340,67 +303,4 @@ class ImportsContextCustomizer implements ContextCustomizer { } - /** - * Filter used to limit considered annotations. - */ - private interface AnnotationFilter { - - boolean isIgnored(Annotation annotation); - - } - - /** - * {@link AnnotationFilter} for {@literal java.lang} annotations. - */ - private static final class JavaLangAnnotationFilter implements AnnotationFilter { - - @Override - public boolean isIgnored(Annotation annotation) { - return AnnotationUtils.isInJavaLangAnnotationPackage(annotation); - } - - } - - /** - * {@link AnnotationFilter} for Kotlin annotations. - */ - private static final class KotlinAnnotationFilter implements AnnotationFilter { - - @Override - public boolean isIgnored(Annotation annotation) { - return "kotlin.Metadata".equals(annotation.annotationType().getName()) - || isInKotlinAnnotationPackage(annotation); - } - - private boolean isInKotlinAnnotationPackage(Annotation annotation) { - return annotation.annotationType().getName().startsWith("kotlin.annotation."); - } - - } - - /** - * {@link AnnotationFilter} for Spock annotations. - */ - private static final class SpockAnnotationFilter implements AnnotationFilter { - - @Override - public boolean isIgnored(Annotation annotation) { - return annotation.annotationType().getName().startsWith("org.spockframework.") - || annotation.annotationType().getName().startsWith("spock."); - } - - } - - /** - * {@link AnnotationFilter} for JUnit annotations. - */ - private static final class JUnitAnnotationFilter implements AnnotationFilter { - - @Override - public boolean isIgnored(Annotation annotation) { - return annotation.annotationType().getName().startsWith("org.junit."); - } - - } - }