Infer proxy on @Lazy-annotated injection points
This commit makes use of the new `getLazyResolutionProxyClass` on `AutowireCandidateResolver` to detect if a injection point requires a proxy. Closes gh-28980
This commit is contained in:
@@ -25,6 +25,7 @@ import java.lang.reflect.InvocationTargetException;
|
||||
import java.lang.reflect.Member;
|
||||
import java.lang.reflect.Method;
|
||||
import java.lang.reflect.Modifier;
|
||||
import java.lang.reflect.Proxy;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
import java.util.Collection;
|
||||
@@ -70,6 +71,8 @@ import org.springframework.beans.factory.config.ConfigurableListableBeanFactory;
|
||||
import org.springframework.beans.factory.config.DependencyDescriptor;
|
||||
import org.springframework.beans.factory.config.SmartInstantiationAwareBeanPostProcessor;
|
||||
import org.springframework.beans.factory.support.AbstractAutowireCapableBeanFactory;
|
||||
import org.springframework.beans.factory.support.AutowireCandidateResolver;
|
||||
import org.springframework.beans.factory.support.DefaultListableBeanFactory;
|
||||
import org.springframework.beans.factory.support.LookupOverride;
|
||||
import org.springframework.beans.factory.support.MergedBeanDefinitionPostProcessor;
|
||||
import org.springframework.beans.factory.support.RegisteredBean;
|
||||
@@ -289,7 +292,7 @@ public class AutowiredAnnotationBeanPostProcessor implements SmartInstantiationA
|
||||
InjectionMetadata metadata = findInjectionMetadata(beanName, beanClass, beanDefinition);
|
||||
Collection<AutowiredElement> autowiredElements = getAutowiredElements(metadata);
|
||||
if (!ObjectUtils.isEmpty(autowiredElements)) {
|
||||
return new AotContribution(beanClass, autowiredElements);
|
||||
return new AotContribution(beanClass, autowiredElements, getAutowireCandidateResolver());
|
||||
}
|
||||
return null;
|
||||
}
|
||||
@@ -300,6 +303,14 @@ public class AutowiredAnnotationBeanPostProcessor implements SmartInstantiationA
|
||||
return (Collection) metadata.getInjectedElements();
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private AutowireCandidateResolver getAutowireCandidateResolver() {
|
||||
if (this.beanFactory instanceof DefaultListableBeanFactory lbf) {
|
||||
return lbf.getAutowireCandidateResolver();
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
private InjectionMetadata findInjectionMetadata(String beanName, Class<?> beanType, RootBeanDefinition beanDefinition) {
|
||||
InjectionMetadata metadata = findAutowiringMetadata(beanName, beanType, null);
|
||||
metadata.checkConfigMembers(beanDefinition);
|
||||
@@ -914,10 +925,15 @@ public class AutowiredAnnotationBeanPostProcessor implements SmartInstantiationA
|
||||
|
||||
private final Collection<AutowiredElement> autowiredElements;
|
||||
|
||||
@Nullable
|
||||
private final AutowireCandidateResolver candidateResolver;
|
||||
|
||||
AotContribution(Class<?> target, Collection<AutowiredElement> autowiredElements,
|
||||
@Nullable AutowireCandidateResolver candidateResolver) {
|
||||
|
||||
AotContribution(Class<?> target, Collection<AutowiredElement> autowiredElements) {
|
||||
this.target = target;
|
||||
this.autowiredElements = autowiredElements;
|
||||
this.candidateResolver = candidateResolver;
|
||||
}
|
||||
|
||||
|
||||
@@ -940,6 +956,10 @@ public class AutowiredAnnotationBeanPostProcessor implements SmartInstantiationA
|
||||
});
|
||||
beanRegistrationCode.addInstancePostProcessor(
|
||||
MethodReference.ofStatic(generatedClass.getName(), generateMethod.getName()));
|
||||
|
||||
if (this.candidateResolver != null) {
|
||||
registerHints(generationContext.getRuntimeHints());
|
||||
}
|
||||
}
|
||||
|
||||
private CodeBlock generateMethodCode(RuntimeHints hints) {
|
||||
@@ -1023,6 +1043,35 @@ public class AutowiredAnnotationBeanPostProcessor implements SmartInstantiationA
|
||||
return code.build();
|
||||
}
|
||||
|
||||
private void registerHints(RuntimeHints runtimeHints) {
|
||||
this.autowiredElements.forEach(autowiredElement -> {
|
||||
boolean required = autowiredElement.required;
|
||||
Member member = autowiredElement.getMember();
|
||||
if (member instanceof Field field) {
|
||||
DependencyDescriptor dependencyDescriptor = new DependencyDescriptor(
|
||||
field, required);
|
||||
registerProxyIfNecessary(runtimeHints, dependencyDescriptor);
|
||||
}
|
||||
if (member instanceof Method method) {
|
||||
Class<?>[] parameterTypes = method.getParameterTypes();
|
||||
for (int i = 0; i < parameterTypes.length; i++) {
|
||||
MethodParameter methodParam = new MethodParameter(method, i);
|
||||
DependencyDescriptor dependencyDescriptor = new DependencyDescriptor(
|
||||
methodParam, required);
|
||||
registerProxyIfNecessary(runtimeHints, dependencyDescriptor);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
private void registerProxyIfNecessary(RuntimeHints runtimeHints, DependencyDescriptor dependencyDescriptor) {
|
||||
Class<?> proxyType = this.candidateResolver
|
||||
.getLazyResolutionProxyClass(dependencyDescriptor, null);
|
||||
if (proxyType != null && Proxy.isProxyClass(proxyType)) {
|
||||
runtimeHints.proxies().registerJdkProxy(proxyType.getInterfaces());
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -16,7 +16,10 @@
|
||||
|
||||
package org.springframework.beans.factory.aot;
|
||||
|
||||
import java.lang.reflect.Constructor;
|
||||
import java.lang.reflect.Executable;
|
||||
import java.lang.reflect.Method;
|
||||
import java.lang.reflect.Proxy;
|
||||
import java.util.List;
|
||||
|
||||
import javax.lang.model.element.Modifier;
|
||||
@@ -26,8 +29,13 @@ import org.springframework.aot.generate.GeneratedMethod;
|
||||
import org.springframework.aot.generate.GeneratedMethods;
|
||||
import org.springframework.aot.generate.GenerationContext;
|
||||
import org.springframework.aot.generate.MethodReference;
|
||||
import org.springframework.aot.hint.RuntimeHints;
|
||||
import org.springframework.beans.factory.config.BeanDefinition;
|
||||
import org.springframework.beans.factory.config.DependencyDescriptor;
|
||||
import org.springframework.beans.factory.support.AutowireCandidateResolver;
|
||||
import org.springframework.beans.factory.support.DefaultListableBeanFactory;
|
||||
import org.springframework.beans.factory.support.RegisteredBean;
|
||||
import org.springframework.core.MethodParameter;
|
||||
import org.springframework.javapoet.ClassName;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.StringUtils;
|
||||
@@ -83,6 +91,7 @@ class BeanDefinitionMethodGenerator {
|
||||
MethodReference generateBeanDefinitionMethod(GenerationContext generationContext,
|
||||
BeanRegistrationsCode beanRegistrationsCode) {
|
||||
|
||||
registerRuntimeHintsIfNecessary(generationContext.getRuntimeHints());
|
||||
BeanRegistrationCodeFragments codeFragments = getCodeFragments(generationContext,
|
||||
beanRegistrationsCode);
|
||||
Class<?> target = codeFragments.getTarget(this.registeredBean,
|
||||
@@ -166,4 +175,54 @@ class BeanDefinitionMethodGenerator {
|
||||
return StringUtils.uncapitalize(beanName);
|
||||
}
|
||||
|
||||
private void registerRuntimeHintsIfNecessary(RuntimeHints runtimeHints) {
|
||||
if (this.registeredBean.getBeanFactory() instanceof DefaultListableBeanFactory dlbf) {
|
||||
ProxyRuntimeHintsRegistrar registrar = new ProxyRuntimeHintsRegistrar(dlbf.getAutowireCandidateResolver());
|
||||
if (this.constructorOrFactoryMethod instanceof Method method) {
|
||||
registrar.registerRuntimeHints(runtimeHints, method);
|
||||
}
|
||||
else if (this.constructorOrFactoryMethod instanceof Constructor<?> constructor) {
|
||||
registrar.registerRuntimeHints(runtimeHints, constructor);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static class ProxyRuntimeHintsRegistrar {
|
||||
|
||||
private final AutowireCandidateResolver candidateResolver;
|
||||
|
||||
public ProxyRuntimeHintsRegistrar(AutowireCandidateResolver candidateResolver) {
|
||||
this.candidateResolver = candidateResolver;
|
||||
}
|
||||
|
||||
public void registerRuntimeHints(RuntimeHints runtimeHints, Method method) {
|
||||
Class<?>[] parameterTypes = method.getParameterTypes();
|
||||
for (int i = 0; i < parameterTypes.length; i++) {
|
||||
MethodParameter methodParam = new MethodParameter(method, i);
|
||||
DependencyDescriptor dependencyDescriptor = new DependencyDescriptor(
|
||||
methodParam, true);
|
||||
registerProxyIfNecessary(runtimeHints, dependencyDescriptor);
|
||||
}
|
||||
}
|
||||
|
||||
public void registerRuntimeHints(RuntimeHints runtimeHints, Constructor<?> constructor) {
|
||||
Class<?>[] parameterTypes = constructor.getParameterTypes();
|
||||
for (int i = 0; i < parameterTypes.length; i++) {
|
||||
MethodParameter methodParam = new MethodParameter(constructor, i);
|
||||
DependencyDescriptor dependencyDescriptor = new DependencyDescriptor(
|
||||
methodParam, true);
|
||||
registerProxyIfNecessary(runtimeHints, dependencyDescriptor);
|
||||
}
|
||||
}
|
||||
|
||||
private void registerProxyIfNecessary(RuntimeHints runtimeHints, DependencyDescriptor dependencyDescriptor) {
|
||||
Class<?> proxyType = this.candidateResolver
|
||||
.getLazyResolutionProxyClass(dependencyDescriptor, null);
|
||||
if (proxyType != null && Proxy.isProxyClass(proxyType)) {
|
||||
runtimeHints.proxies().registerJdkProxy(proxyType.getInterfaces());
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user