diff --git a/spring-test/src/main/java/org/springframework/test/context/bean/override/BeanOverrideUtils.java b/spring-test/src/main/java/org/springframework/test/context/bean/override/BeanOverrideUtils.java index c97e1b2c91b..a83b56e7588 100644 --- a/spring-test/src/main/java/org/springframework/test/context/bean/override/BeanOverrideUtils.java +++ b/spring-test/src/main/java/org/springframework/test/context/bean/override/BeanOverrideUtils.java @@ -29,6 +29,7 @@ import java.util.HashSet; import java.util.List; import java.util.Set; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiConsumer; import kotlin.jvm.JvmClassMappingKt; @@ -62,6 +63,30 @@ public abstract class BeanOverrideUtils { Comparator.> comparingInt(MergedAnnotation::getDistance).reversed(); + /** + * Resolve the {@link BeanOverrideHandler} for the given {@link Parameter}. + * @param parameter the parameter to process + * @param testClass the test class to process + * @return the bean override handler for the parameter, or {@code null} if no + * handler was found + * @see BeanOverrideProcessor#createHandler(Annotation, Class, Parameter) + * @see #findAllHandlers(Class) + */ + public static @Nullable BeanOverrideHandler resolveHandlerForParameter(Parameter parameter, Class testClass) { + AtomicReference handlerReference = new AtomicReference<>(); + AtomicBoolean overrideAnnotationFound = new AtomicBoolean(); + processElement(parameter, (processor, composedAnnotation) -> { + Assert.state(overrideAnnotationFound.compareAndSet(false, true), + () -> "Multiple @BeanOverride annotations found on parameter: " + parameter); + BeanOverrideHandler handler = processor.createHandler(composedAnnotation, testClass, parameter); + Assert.state(handler != null, + () -> "BeanOverrideProcessor [%s] returned null BeanOverrideHandler for parameter [%s]" + .formatted(processor.getClass().getSimpleName(), parameter)); + handlerReference.setPlain(handler); + }); + return handlerReference.getPlain(); + } + /** * Process the given {@code testClass} and build the corresponding * {@link BeanOverrideHandler} list derived from {@link BeanOverride @BeanOverride} @@ -167,16 +192,10 @@ public abstract class BeanOverrideUtils { } private static void processParameter(Parameter parameter, Class testClass, List handlers) { - AtomicBoolean overrideAnnotationFound = new AtomicBoolean(); - processElement(parameter, (processor, composedAnnotation) -> { - Assert.state(overrideAnnotationFound.compareAndSet(false, true), - () -> "Multiple @BeanOverride annotations found on parameter: " + parameter); - BeanOverrideHandler handler = processor.createHandler(composedAnnotation, testClass, parameter); - Assert.state(handler != null, - () -> "BeanOverrideProcessor [%s] returned null BeanOverrideHandler for parameter [%s]" - .formatted(processor.getClass().getSimpleName(), parameter)); + BeanOverrideHandler handler = resolveHandlerForParameter(parameter, testClass); + if (handler != null) { handlers.add(handler); - }); + } } private static void processField(Field field, Class testClass, List handlers) { diff --git a/spring-test/src/main/java/org/springframework/test/context/junit/jupiter/SpringExtension.java b/spring-test/src/main/java/org/springframework/test/context/junit/jupiter/SpringExtension.java index 92f186c7804..26bccfaee4b 100644 --- a/spring-test/src/main/java/org/springframework/test/context/junit/jupiter/SpringExtension.java +++ b/spring-test/src/main/java/org/springframework/test/context/junit/jupiter/SpringExtension.java @@ -24,8 +24,6 @@ import java.lang.reflect.Parameter; import java.util.Arrays; import java.util.List; import java.util.Locale; -import java.util.Objects; -import java.util.Optional; import org.jspecify.annotations.Nullable; import org.junit.jupiter.api.AfterAll; @@ -437,14 +435,13 @@ public class SpringExtension implements BeforeAllCallback, AfterAllCallback, Tes } ApplicationContext applicationContext = getApplicationContext(extensionContext); + + // If the parameter is a @BeanOverride with an explicit name, we simply look + // up the bean by name instead of performing full dependency resolution. if (isBeanOverride(parameter)) { - Optional beanName = BeanOverrideUtils.findAllHandlers(testClass).stream() - .filter(handler -> parameter.equals(handler.getParameter())) - .map(BeanOverrideHandler::getBeanName) - .filter(Objects::nonNull) - .findFirst(); - if (beanName.isPresent()) { - return applicationContext.getBean(beanName.get()); + BeanOverrideHandler handler = BeanOverrideUtils.resolveHandlerForParameter(parameter, testClass); + if (handler != null && handler.getBeanName() != null) { + return applicationContext.getBean(handler.getBeanName()); } } return ParameterResolutionDelegate.resolveDependency(parameter, index, testClass,