From ce15b8e5802ad435ab296693ac84dd41d64106f1 Mon Sep 17 00:00:00 2001 From: Mehmet Mustafa Yilmaz Date: Wed, 12 Oct 2022 17:36:05 -0700 Subject: [PATCH] fix: Guice Cannot Inject Beans with Custom Annotations SpringModule binds custom guice Providers for Spring managed beans so that Guice can inject beans from Spring Context. Due to a bug in the Provider, Guice can only inject Spring beans that either don't have any qualifier annotations or only have Named annotation as a qualifier. This fix enables Guice to inject beans with custom qualifier annotations as well. Custom qualifier annotations do not need to be marker annotations (in other words, they can have attributes). Other changes include: Using factory method metadata of annotated bean definition rather than using custom code to retrieve factory method and its annotations Added more test cases to validate various qualifier annotation scenarios. --- .../guice/module/SpringModule.java | 137 ++++++++---------- .../module/SpringModuleMetadataTests.java | 125 +++++++++++++++- 2 files changed, 182 insertions(+), 80 deletions(-) diff --git a/src/main/java/org/springframework/guice/module/SpringModule.java b/src/main/java/org/springframework/guice/module/SpringModule.java index f373d59..d4d9f45 100644 --- a/src/main/java/org/springframework/guice/module/SpringModule.java +++ b/src/main/java/org/springframework/guice/module/SpringModule.java @@ -17,11 +17,9 @@ package org.springframework.guice.module; import java.lang.annotation.Annotation; -import java.lang.reflect.Method; import java.lang.reflect.ParameterizedType; import java.lang.reflect.Type; import java.util.ArrayList; -import java.util.Arrays; import java.util.Collection; import java.util.HashMap; import java.util.HashSet; @@ -58,11 +56,9 @@ import org.springframework.beans.factory.support.DefaultListableBeanFactory; import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.context.ApplicationContext; import org.springframework.core.ResolvableType; -import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.core.annotation.MergedAnnotation; import org.springframework.core.type.MethodMetadata; -import org.springframework.core.type.StandardMethodMetadata; import org.springframework.util.ClassUtils; -import org.springframework.util.ReflectionUtils; /** * A Guice module that wraps a Spring {@link ApplicationContext}. @@ -137,7 +133,7 @@ public class SpringModule extends AbstractModule { if (definition.hasAttribute(SPRING_GUICE_SOURCE)) { continue; } - Optional bindingAnnotation = getAnnotationForBeanDefinition(definition, beanFactory); + Optional bindingAnnotation = getAnnotationForBeanDefinition(definition); if (definition.isAutowireCandidate() && definition.getRole() == AbstractBeanDefinition.ROLE_APPLICATION) { Type type; Class clazz = beanFactory.getType(name); @@ -204,16 +200,15 @@ public class SpringModule extends AbstractModule { } } - private static Optional getAnnotationForBeanDefinition(BeanDefinition definition, - ConfigurableListableBeanFactory beanFactory) { - if (definition instanceof AnnotatedBeanDefinition - && ((AnnotatedBeanDefinition) definition).getFactoryMethodMetadata() != null) { - try { - Method factoryMethod = getFactoryMethod(beanFactory, definition); - return Arrays.stream(AnnotationUtils.getAnnotations(factoryMethod)) - .filter((a) -> Annotations.isBindingAnnotation(a.annotationType())).findFirst(); + private static Optional getAnnotationForBeanDefinition(BeanDefinition definition) { + if (definition instanceof AnnotatedBeanDefinition) { + MethodMetadata methodMetadata = ((AnnotatedBeanDefinition) definition).getFactoryMethodMetadata(); + if (methodMetadata != null) { + return methodMetadata.getAnnotations().stream().filter(MergedAnnotation::isDirectlyPresent) + .filter((mergedAnnotation) -> Annotations.isBindingAnnotation(mergedAnnotation.getType())) + .map(MergedAnnotation::synthesize).findFirst(); } - catch (Exception ex) { + else { return Optional.empty(); } } @@ -222,49 +217,6 @@ public class SpringModule extends AbstractModule { } } - private static Method getFactoryMethod(ConfigurableListableBeanFactory beanFactory, BeanDefinition definition) - throws Exception { - if (definition instanceof AnnotatedBeanDefinition) { - MethodMetadata factoryMethodMetadata = ((AnnotatedBeanDefinition) definition).getFactoryMethodMetadata(); - if (factoryMethodMetadata instanceof StandardMethodMetadata) { - return ((StandardMethodMetadata) factoryMethodMetadata).getIntrospectedMethod(); - } - } - BeanDefinition factoryDefinition = beanFactory.getBeanDefinition(definition.getFactoryBeanName()); - Class factoryClass = ClassUtils.forName(factoryDefinition.getBeanClassName(), - beanFactory.getBeanClassLoader()); - return getFactoryMethod(definition, factoryClass); - } - - private static Method getFactoryMethod(BeanDefinition definition, Class factoryClass) { - Method uniqueMethod = null; - for (Method candidate : getCandidateFactoryMethods(definition, factoryClass)) { - if (candidate.getName().equals(definition.getFactoryMethodName())) { - if (uniqueMethod == null) { - uniqueMethod = candidate; - } - else if (!hasMatchingParameterTypes(candidate, uniqueMethod)) { - return null; - } - } - } - return uniqueMethod; - } - - private static Method[] getCandidateFactoryMethods(BeanDefinition definition, Class factoryClass) { - return shouldConsiderNonPublicMethods(definition) ? ReflectionUtils.getAllDeclaredMethods(factoryClass) - : factoryClass.getMethods(); - } - - private static boolean shouldConsiderNonPublicMethods(BeanDefinition definition) { - return (definition instanceof AbstractBeanDefinition) - && ((AbstractBeanDefinition) definition).isNonPublicAccessAllowed(); - } - - private static boolean hasMatchingParameterTypes(Method candidate, Method current) { - return Arrays.equals(candidate.getParameterTypes(), current.getParameterTypes()); - } - private static Set getAllSuperTypes(Type originalType, Class clazz) { Set allInterfaces = new HashSet<>(); TypeLiteral typeToken = TypeLiteral.get(originalType); @@ -420,34 +372,65 @@ public class SpringModule extends AbstractModule { String[] named = BeanFactoryUtils.beanNamesForTypeIncludingAncestors(this.beanFactory, ResolvableType.forType(this.type)); - List names = new ArrayList(named.length); - if (named.length == 1) { - names.add(named[0]); + + List candidateBeanNames = new ArrayList<>(named.length); + for (String name : named) { + BeanDefinition beanDefinition = this.beanFactory.getBeanDefinition(name); + // This is a Guice component bridged to spring + // If this were the target candidate, + // Guice would have injected it natively. + // Thus, it cannot be a candidate. + // GuiceFactoryBeans don't have 1-to-1 annotation mapping + // (since annotation attributes are ignored) + // Skip this candidate to avoid unexpected matches + // due to imprecise annotation mapping + if (!beanDefinition.hasAttribute(SPRING_GUICE_SOURCE)) { + candidateBeanNames.add(name); + } + } + + List matchingBeanNames; + if (candidateBeanNames.size() == 1) { + matchingBeanNames = candidateBeanNames; } else { - for (String name : named) { - if (this.bindingAnnotation.isPresent()) { - if (this.bindingAnnotation.get() instanceof Named - || this.bindingAnnotation.get() instanceof javax.inject.Named) { - Optional annotation = SpringModule.getAnnotationForBeanDefinition( - this.beanFactory.getMergedBeanDefinition(name), this.beanFactory); - String boundName = getNameFromBindingAnnotation(this.bindingAnnotation); - if (annotation.isPresent() && this.bindingAnnotation.get().equals(annotation.get()) - || name.equals(boundName)) { - names.add(name); + matchingBeanNames = new ArrayList(candidateBeanNames.size()); + for (String name : candidateBeanNames) { + // Make sure we don't add the same name twice using if/else + if (name.equals(this.name)) { + // Guice is injecting dependency of this type by bean name + matchingBeanNames.add(name); + } + else if (this.bindingAnnotation.isPresent()) { + String boundName = getNameFromBindingAnnotation(this.bindingAnnotation); + if (name.equals(boundName)) { + // Spring bean definition has a Named annotation that + // matches the name of the bean + // In such cases, we dedupe namedProvider (because it's + // Key equals typeProvider Key) + // Thus, this complementary check is required + // (because name field is null in typeProvider, + // and if check above wouldn't pass) + matchingBeanNames.add(name); + } + else { + Optional annotationOptional = SpringModule + .getAnnotationForBeanDefinition(this.beanFactory.getBeanDefinition(name)); + + if (annotationOptional.equals(this.bindingAnnotation)) { + // Found a bean with matching qualifier annotation + matchingBeanNames.add(name); } } } - if (name.equals(this.name)) { - names.add(name); - } } } - if (names.size() == 1) { - this.resultProvider = () -> this.beanFactory.getBean(names.get(0)); + if (matchingBeanNames.size() == 1) { + this.resultProvider = () -> this.beanFactory.getBean(matchingBeanNames.get(0)); } else { - for (String name : named) { + // Shouldn't we iterate over matching bean names here? + for (String name : candidateBeanNames) { if (this.beanFactory.getBeanDefinition(name).isPrimary()) { this.resultProvider = () -> this.beanFactory.getBean(name); break; diff --git a/src/test/java/org/springframework/guice/module/SpringModuleMetadataTests.java b/src/test/java/org/springframework/guice/module/SpringModuleMetadataTests.java index a091776..cf94bcd 100644 --- a/src/test/java/org/springframework/guice/module/SpringModuleMetadataTests.java +++ b/src/test/java/org/springframework/guice/module/SpringModuleMetadataTests.java @@ -16,7 +16,12 @@ package org.springframework.guice.module; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; + import javax.inject.Inject; +import javax.inject.Named; +import javax.inject.Qualifier; import com.google.inject.ConfigurationException; import com.google.inject.Guice; @@ -35,6 +40,7 @@ import org.springframework.core.type.filter.AnnotationTypeFilter; import org.springframework.core.type.filter.AssignableTypeFilter; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatCode; import static org.assertj.core.api.Assertions.assertThatExceptionOfType; /** @@ -68,6 +74,43 @@ public class SpringModuleMetadataTests { assertThat(injector.getInstance(Key.get(Service.class, Names.named("service")))).isNotNull(); } + @Test + public void threeServicesByQualifier() throws Exception { + Injector injector = createInjector(PrimaryConfig.class, QualifiedConfig.class); + + assertThat(injector.getInstance( + Key.get(Service.class, ServiceQualifierAnnotated.class.getAnnotation(ServiceQualifier.class)))) + .extracting("name").isEqualTo("emptyQualifierService"); + + assertThat(injector.getInstance( + Key.get(Service.class, EmptyServiceQualifierAnnotated.class.getAnnotation(ServiceQualifier.class)))) + .extracting("name").isEqualTo("emptyQualifierService"); + + assertThat(injector.getInstance( + Key.get(Service.class, MyServiceQualifierAnnotated.class.getAnnotation(ServiceQualifier.class)))) + .extracting("name").isEqualTo("myService"); + + assertThat(injector.getInstance(Key.get(Service.class, Names.named("namedService")))).extracting("name") + .isEqualTo("namedService"); + + assertThat(injector.getInstance(Key.get(Service.class, Names.named("namedServiceWithADifferentBeanName")))) + .extracting("name").isEqualTo("namedServiceWithADifferentBeanName"); + + assertThat(injector.getInstance(Service.class)).extracting("name").isEqualTo("primary"); + + // Test cases where we don't expect to find any bindings + assertThatCode(() -> injector.getInstance(Key.get(Service.class, Names.named("randomService")))) + .isInstanceOf(ConfigurationException.class); + + assertThatCode(() -> injector.getInstance( + Key.get(Service.class, NoServiceQualifierAnnotated.class.getAnnotation(ServiceQualifier.class)))) + .isInstanceOf(ConfigurationException.class); + + assertThatCode(() -> injector.getInstance(Key.get(Service.class, UnboundServiceQualifier.class))) + .isInstanceOf(ConfigurationException.class); + + } + @Test public void includes() throws Exception { Injector injector = createInjector(TestConfig.class, MetadataIncludesConfig.class); @@ -92,10 +135,23 @@ public class SpringModuleMetadataTests { interface Service { + String getName(); + } protected static class MyService implements Service { + private final String name; + + protected MyService(String name) { + this.name = name; + } + + @Override + public String getName() { + return this.name; + } + } public static class Foo { @@ -135,7 +191,7 @@ public class SpringModuleMetadataTests { @Bean public Service service() { - return new MyService(); + return new MyService("service"); } } @@ -146,7 +202,7 @@ public class SpringModuleMetadataTests { @Bean @Primary public Service primary() { - return new MyService(); + return new MyService("primary"); } } @@ -156,7 +212,36 @@ public class SpringModuleMetadataTests { @Bean public Service more() { - return new MyService(); + return new MyService("more"); + } + + } + + @Configuration + public static class QualifiedConfig { + + @Bean + @Named("namedService") + public Service namedService() { + return new MyService("namedService"); + } + + @Bean + @Named("namedServiceWithADifferentBeanName") + public Service anotherNamedService() { + return new MyService("namedServiceWithADifferentBeanName"); + } + + @Bean + @ServiceQualifier + public Service emptyQualifierService() { + return new MyService("emptyQualifierService"); + } + + @Bean + @ServiceQualifier(type = "myService") + public Service myService(@Named("namedService") Service service) { + return new MyService("myService"); } } @@ -166,4 +251,38 @@ public class SpringModuleMetadataTests { } + @Qualifier + @Retention(RetentionPolicy.RUNTIME) + @interface ServiceQualifier { + + String type() default ""; + + } + + @Qualifier + @Retention(RetentionPolicy.RUNTIME) + @interface UnboundServiceQualifier { + + } + + @ServiceQualifier + interface ServiceQualifierAnnotated { + + } + + @ServiceQualifier(type = "") + interface EmptyServiceQualifierAnnotated { + + } + + @ServiceQualifier(type = "myService") + interface MyServiceQualifierAnnotated { + + } + + @ServiceQualifier(type = "noService") + interface NoServiceQualifierAnnotated { + + } + }