From 74e52612bf2ee6031b2a1a12b881e00402743f90 Mon Sep 17 00:00:00 2001 From: Christoph Strobl Date: Wed, 13 Jul 2022 10:57:18 +0200 Subject: [PATCH] Use ManagedType BeanDefinition for AOT processing when possible. We now try to read the types directly from the bean definition arguments first, before attempting to resolve the actual bean instance. If resolving the bean fails, we currently only log an info message. This arrangement needs to be revisited. See: #2593 --- ...agedTypesBeanRegistrationAotProcessor.java | 42 +++++++++++++++-- ...BeanRegistrationAotProcessorUnitTests.java | 45 ++++++++++++++++++- 2 files changed, 83 insertions(+), 4 deletions(-) diff --git a/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java b/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java index 3353d1d82..77c30babe 100644 --- a/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java +++ b/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java @@ -15,17 +15,20 @@ */ package org.springframework.data.aot; +import java.util.Collection; import java.util.Collections; import java.util.Set; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; - import org.springframework.aot.generate.GenerationContext; +import org.springframework.beans.factory.BeanCreationException; import org.springframework.beans.factory.BeanFactory; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; import org.springframework.beans.factory.aot.BeanRegistrationAotProcessor; +import org.springframework.beans.factory.config.ConstructorArgumentValues.ValueHolder; import org.springframework.beans.factory.support.RegisteredBean; +import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.core.ResolvableType; import org.springframework.data.domain.ManagedTypes; import org.springframework.lang.Nullable; @@ -53,8 +56,41 @@ public class ManagedTypesBeanRegistrationAotProcessor implements BeanRegistratio } BeanFactory beanFactory = registeredBean.getBeanFactory(); - return contribute(AotContext.from(registeredBean.getBeanFactory()), - beanFactory.getBean(registeredBean.getBeanName(), ManagedTypes.class)); + return contribute(AotContext.from(beanFactory), resolveManagedTypes(registeredBean)); + } + + ManagedTypes resolveManagedTypes(RegisteredBean registeredBean) { + + RootBeanDefinition beanDefinition = registeredBean.getMergedBeanDefinition(); + if (beanDefinition.hasConstructorArgumentValues()) { + ValueHolder indexedArgumentValue = beanDefinition.getConstructorArgumentValues().getIndexedArgumentValue(0, null); + Object value = indexedArgumentValue.getValue(); + if (value instanceof Collection values) { + if (values.stream().allMatch(it -> it instanceof Class)) { + return ManagedTypes.fromIterable((Collection>) values); + } + } + } + if (logger.isDebugEnabled()) { + logger.debug( + String.format("ManagedTypes BeanDefinition '%s' does serve arguments. Trying to resolve bean instance.", + registeredBean.getBeanName())); + } + + if (registeredBean.getParent() == null) { + try { + return registeredBean.getBeanFactory().getBean(registeredBean.getBeanName(), ManagedTypes.class); + } catch (BeanCreationException e) { + if (logger.isInfoEnabled()) { + logger.info(String.format("Could not resolve ManagedTypes '%s'.", registeredBean.getBeanName())); + } + if (logger.isDebugEnabled()) { + logger.debug(e); + } + } + } + + return ManagedTypes.empty(); } protected boolean isMatch(@Nullable Class beanType, @Nullable String beanName) { diff --git a/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java b/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java index 8d5d50e0e..24adc95b2 100644 --- a/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java +++ b/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java @@ -16,6 +16,8 @@ package org.springframework.data.aot; import static org.assertj.core.api.Assertions.*; +import static org.mockito.ArgumentMatchers.*; +import static org.mockito.Mockito.*; import java.util.Collections; import java.util.function.Consumer; @@ -28,6 +30,7 @@ import org.springframework.aot.generate.GeneratedClasses; import org.springframework.aot.generate.InMemoryGeneratedFiles; import org.springframework.aot.hint.RuntimeHints; import org.springframework.aot.hint.predicate.RuntimeHintsPredicates; +import org.springframework.beans.factory.BeanCreationException; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; import org.springframework.beans.factory.support.BeanDefinitionBuilder; import org.springframework.beans.factory.support.DefaultListableBeanFactory; @@ -51,7 +54,7 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { @BeforeEach void beforeEach() { - beanFactory = new DefaultListableBeanFactory(); + beanFactory = spy(new DefaultListableBeanFactory()); } @Test // GH-2593 @@ -65,6 +68,16 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { assertThat(contribution).isNotNull(); } + @Test // GH-2593 + void processesBeanDefinitionIfPossibleWithoutLoadingTheBean() { + + beanFactory.registerBeanDefinition("commons.managed-types", managedTypesDefinition); + + createPostProcessor("commons").processAheadOfTime(RegisteredBean.of(beanFactory, "commons.managed-types")); + + verify(beanFactory, never()).getBean(eq("commons.managed-types"), eq(ManagedTypes.class)); + } + @Test // GH-2593 void contributesReflectionForManagedTypes() { @@ -94,6 +107,16 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { assertThat(contribution).isNotNull(); } + @Test // GH-2593 + void processesMatchingSubtypeBeanByAttemptingToLoadItIfNoMatchingConstructorArgumentFound() { + + beanFactory.registerBeanDefinition("commons.managed-types", myManagedTypesDefinition); + + createPostProcessor("commons").processAheadOfTime(RegisteredBean.of(beanFactory, "commons.managed-types")); + + verify(beanFactory).getBean(eq("commons.managed-types"), eq(ManagedTypes.class)); + } + @Test // GH-2593 void ignoresBeanNotMatchingRequiredType() { @@ -117,6 +140,26 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { assertThat(contribution).isNull(); } + @Test // GH-2593 + void returnsEmptyContributionWhenBeanCannotBeLoaded() { + + doThrow(new BeanCreationException("o_O")).when(beanFactory).getBean(eq("commons.managed-types"), + eq(ManagedTypes.class)); + + beanFactory.registerBeanDefinition("commons.managed-types", myManagedTypesDefinition); + + BeanRegistrationAotContribution contribution = createPostProcessor("commons") + .processAheadOfTime(RegisteredBean.of(beanFactory, "commons.managed-types")); + + DefaultGenerationContext generationContext = new DefaultGenerationContext( + new GeneratedClasses(new ClassNameGenerator(Object.class)), new InMemoryGeneratedFiles(), new RuntimeHints()); + + contribution.applyTo(generationContext, null); + + assertThat(generationContext.getRuntimeHints().reflection().typeHints()).isEmpty(); + verify(beanFactory).getBean(eq("commons.managed-types"), eq(ManagedTypes.class)); + } + private ManagedTypesBeanRegistrationAotProcessor createPostProcessor(String moduleIdentifier) { ManagedTypesBeanRegistrationAotProcessor postProcessor = new ManagedTypesBeanRegistrationAotProcessor(); postProcessor.setModuleIdentifier(moduleIdentifier);