diff --git a/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContribution.java b/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContribution.java index a80db112d3..d15191a3bb 100644 --- a/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContribution.java +++ b/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContribution.java @@ -26,7 +26,10 @@ import org.springframework.aot.generate.GeneratedMethods; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.generate.MethodReference; import org.springframework.aot.generate.MethodReference.ArgumentCodeGenerator; +import org.springframework.aot.hint.MemberCategory; +import org.springframework.aot.hint.RuntimeHints; import org.springframework.beans.factory.support.DefaultListableBeanFactory; +import org.springframework.beans.factory.support.RegisteredBean; import org.springframework.javapoet.ClassName; import org.springframework.javapoet.CodeBlock; import org.springframework.javapoet.MethodSpec; @@ -36,6 +39,7 @@ import org.springframework.javapoet.MethodSpec; * register bean definitions. * * @author Phillip Webb + * @author Brian Clozel * @since 6.0 * @see BeanRegistrationsAotProcessor */ @@ -44,11 +48,11 @@ class BeanRegistrationsAotContribution private static final String BEAN_FACTORY_PARAMETER_NAME = "beanFactory"; - private final Map registrations; + private final Map registrations; BeanRegistrationsAotContribution( - Map registrations) { + Map registrations) { this.registrations = registrations; } @@ -67,8 +71,15 @@ class BeanRegistrationsAotContribution GeneratedMethod generatedMethod = codeGenerator.getMethods().add("registerBeanDefinitions", method -> generateRegisterMethod(method, generationContext, codeGenerator)); beanFactoryInitializationCode.addInitializer(generatedMethod.toMethodReference()); + generateRegisterHints(generationContext.getRuntimeHints(), this.registrations); } + private void generateRegisterHints(RuntimeHints runtimeHints, Map registrations) { + registrations.keySet().forEach(registeredBean -> runtimeHints.reflection() + .registerType(registeredBean.getBeanClass(), MemberCategory.INTROSPECT_DECLARED_METHODS)); + } + + private void generateRegisterMethod(MethodSpec.Builder method, GenerationContext generationContext, BeanRegistrationsCode beanRegistrationsCode) { @@ -78,14 +89,14 @@ class BeanRegistrationsAotContribution method.addParameter(DefaultListableBeanFactory.class, BEAN_FACTORY_PARAMETER_NAME); CodeBlock.Builder code = CodeBlock.builder(); - this.registrations.forEach((beanName, beanDefinitionMethodGenerator) -> { + this.registrations.forEach((registeredBean, beanDefinitionMethodGenerator) -> { MethodReference beanDefinitionMethod = beanDefinitionMethodGenerator .generateBeanDefinitionMethod(generationContext, beanRegistrationsCode); CodeBlock methodInvocation = beanDefinitionMethod.toInvokeCodeBlock( ArgumentCodeGenerator.none(), beanRegistrationsCode.getClassName()); code.addStatement("$L.registerBeanDefinition($S, $L)", - BEAN_FACTORY_PARAMETER_NAME, beanName, + BEAN_FACTORY_PARAMETER_NAME, registeredBean.getBeanName(), methodInvocation); }); method.addCode(code.build()); diff --git a/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessor.java b/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessor.java index 9388d79919..50545dfbdf 100644 --- a/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessor.java +++ b/spring-beans/src/main/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessor.java @@ -27,6 +27,7 @@ import org.springframework.beans.factory.support.RegisteredBean; * register beans. * * @author Phillip Webb + * @author Brian Clozel * @since 6.0 */ class BeanRegistrationsAotProcessor implements BeanFactoryInitializationAotProcessor { @@ -35,13 +36,13 @@ class BeanRegistrationsAotProcessor implements BeanFactoryInitializationAotProce public BeanRegistrationsAotContribution processAheadOfTime(ConfigurableListableBeanFactory beanFactory) { BeanDefinitionMethodGeneratorFactory beanDefinitionMethodGeneratorFactory = new BeanDefinitionMethodGeneratorFactory(beanFactory); - Map registrations = new LinkedHashMap<>(); + Map registrations = new LinkedHashMap<>(); for (String beanName : beanFactory.getBeanDefinitionNames()) { RegisteredBean registeredBean = RegisteredBean.of(beanFactory, beanName); BeanDefinitionMethodGenerator beanDefinitionMethodGenerator = beanDefinitionMethodGeneratorFactory .getBeanDefinitionMethodGenerator(registeredBean, null); if (beanDefinitionMethodGenerator != null) { - registrations.put(beanName, beanDefinitionMethodGenerator); + registrations.put(registeredBean, beanDefinitionMethodGenerator); } } if (registrations.isEmpty()) { diff --git a/spring-beans/src/main/java/org/springframework/beans/factory/support/RegisteredBean.java b/spring-beans/src/main/java/org/springframework/beans/factory/support/RegisteredBean.java index d4ff6f070e..5e93e5ac9e 100644 --- a/spring-beans/src/main/java/org/springframework/beans/factory/support/RegisteredBean.java +++ b/spring-beans/src/main/java/org/springframework/beans/factory/support/RegisteredBean.java @@ -16,6 +16,7 @@ package org.springframework.beans.factory.support; +import java.util.Objects; import java.util.function.BiFunction; import java.util.function.Supplier; @@ -195,6 +196,23 @@ public final class RegisteredBean { return this.parent; } + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (o == null || getClass() != o.getClass()) { + return false; + } + RegisteredBean that = (RegisteredBean) o; + return this.beanName.equals(that.beanName); + } + + @Override + public int hashCode() { + return Objects.hash(this.beanName); + } + @Override public String toString() { return new ToStringCreator(this).append("beanName", getBeanName()) diff --git a/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContributionTests.java b/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContributionTests.java index 577631ff54..c059540591 100644 --- a/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContributionTests.java +++ b/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotContributionTests.java @@ -32,6 +32,8 @@ import org.springframework.aot.generate.ClassNameGenerator; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.generate.MethodReference; import org.springframework.aot.generate.MethodReference.ArgumentCodeGenerator; +import org.springframework.aot.hint.MemberCategory; +import org.springframework.aot.hint.predicate.RuntimeHintsPredicates; import org.springframework.aot.test.generate.TestGenerationContext; import org.springframework.beans.factory.support.DefaultListableBeanFactory; import org.springframework.beans.factory.support.RegisteredBean; @@ -76,13 +78,13 @@ class BeanRegistrationsAotContributionTests { @Test void applyToAppliesContribution() { - Map registrations = new LinkedHashMap<>(); + Map registrations = new LinkedHashMap<>(); RegisteredBean registeredBean = registerBean( new RootBeanDefinition(TestBean.class)); BeanDefinitionMethodGenerator generator = new BeanDefinitionMethodGenerator( this.methodGeneratorFactory, registeredBean, null, Collections.emptyList()); - registrations.put("testBean", generator); + registrations.put(registeredBean, generator); BeanRegistrationsAotContribution contribution = new BeanRegistrationsAotContribution( registrations); contribution.applyTo(this.generationContext, this.beanFactoryInitializationCode); @@ -98,13 +100,13 @@ class BeanRegistrationsAotContributionTests { this.generationContext = new TestGenerationContext( new ClassNameGenerator(TestGenerationContext.TEST_TARGET, "Management")); this.beanFactoryInitializationCode = new MockBeanFactoryInitializationCode(this.generationContext); - Map registrations = new LinkedHashMap<>(); + Map registrations = new LinkedHashMap<>(); RegisteredBean registeredBean = registerBean( new RootBeanDefinition(TestBean.class)); BeanDefinitionMethodGenerator generator = new BeanDefinitionMethodGenerator( this.methodGeneratorFactory, registeredBean, null, Collections.emptyList()); - registrations.put("testBean", generator); + registrations.put(registeredBean, generator); BeanRegistrationsAotContribution contribution = new BeanRegistrationsAotContribution( registrations); contribution.applyTo(this.generationContext, this.beanFactoryInitializationCode); @@ -117,7 +119,7 @@ class BeanRegistrationsAotContributionTests { @Test void applyToCallsRegistrationsWithBeanRegistrationsCode() { List beanRegistrationsCodes = new ArrayList<>(); - Map registrations = new LinkedHashMap<>(); + Map registrations = new LinkedHashMap<>(); RegisteredBean registeredBean = registerBean( new RootBeanDefinition(TestBean.class)); BeanDefinitionMethodGenerator generator = new BeanDefinitionMethodGenerator( @@ -134,7 +136,7 @@ class BeanRegistrationsAotContributionTests { } }; - registrations.put("testBean", generator); + registrations.put(registeredBean, generator); BeanRegistrationsAotContribution contribution = new BeanRegistrationsAotContribution( registrations); contribution.applyTo(this.generationContext, this.beanFactoryInitializationCode); @@ -143,6 +145,22 @@ class BeanRegistrationsAotContributionTests { assertThat(actual.getMethods()).isNotNull(); } + @Test + void applyToRegisterReflectionHints() { + Map registrations = new LinkedHashMap<>(); + RegisteredBean registeredBean = registerBean( + new RootBeanDefinition(TestBean.class)); + BeanDefinitionMethodGenerator generator = new BeanDefinitionMethodGenerator( + this.methodGeneratorFactory, registeredBean, null, + Collections.emptyList()); + registrations.put(registeredBean, generator); + BeanRegistrationsAotContribution contribution = new BeanRegistrationsAotContribution( + registrations); + contribution.applyTo(this.generationContext, this.beanFactoryInitializationCode); + assertThat(RuntimeHintsPredicates.reflection().onType(TestBean.class).withMemberCategory(MemberCategory.INTROSPECT_DECLARED_METHODS)) + .accepts(this.generationContext.getRuntimeHints()); + } + private RegisteredBean registerBean(RootBeanDefinition rootBeanDefinition) { String beanName = "testBean"; this.beanFactory.registerBeanDefinition(beanName, rootBeanDefinition); diff --git a/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessorTests.java b/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessorTests.java index efd8fcdbda..9586c536b9 100644 --- a/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessorTests.java +++ b/spring-beans/src/test/java/org/springframework/beans/factory/aot/BeanRegistrationsAotProcessorTests.java @@ -50,7 +50,7 @@ class BeanRegistrationsAotProcessorTests { BeanRegistrationsAotContribution contribution = processor .processAheadOfTime(beanFactory); assertThat(contribution).extracting("registrations") - .asInstanceOf(InstanceOfAssertFactories.MAP).containsKeys("b1", "b2"); + .asInstanceOf(InstanceOfAssertFactories.MAP).hasSize(2); } }