diff --git a/pom.xml b/pom.xml index 29ed32f2a..f637fcad1 100644 --- a/pom.xml +++ b/pom.xml @@ -121,6 +121,11 @@ reactor-test test + + org.springframework + spring-core-test + test + diff --git a/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java b/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java index 77c30babe..4bd4fe92e 100644 --- a/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java +++ b/src/main/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessor.java @@ -56,7 +56,7 @@ public class ManagedTypesBeanRegistrationAotProcessor implements BeanRegistratio } BeanFactory beanFactory = registeredBean.getBeanFactory(); - return contribute(AotContext.from(beanFactory), resolveManagedTypes(registeredBean)); + return contribute(AotContext.from(beanFactory), resolveManagedTypes(registeredBean), registeredBean); } ManagedTypes resolveManagedTypes(RegisteredBean registeredBean) { @@ -114,8 +114,8 @@ public class ManagedTypesBeanRegistrationAotProcessor implements BeanRegistratio * @return new instance of {@link ManagedTypesBeanRegistrationAotProcessor} or {@literal null} if nothing to do. */ @Nullable - protected BeanRegistrationAotContribution contribute(AotContext aotContext, ManagedTypes managedTypes) { - return new ManagedTypesRegistrationAotContribution(aotContext, managedTypes, this::contributeType); + protected BeanRegistrationAotContribution contribute(AotContext aotContext, ManagedTypes managedTypes, RegisteredBean registeredBean) { + return new ManagedTypesRegistrationAotContribution(aotContext, managedTypes, registeredBean, this::contributeType); } /** diff --git a/src/main/java/org/springframework/data/aot/ManagedTypesRegistrationAotContribution.java b/src/main/java/org/springframework/data/aot/ManagedTypesRegistrationAotContribution.java index 3b81186a2..da455df23 100644 --- a/src/main/java/org/springframework/data/aot/ManagedTypesRegistrationAotContribution.java +++ b/src/main/java/org/springframework/data/aot/ManagedTypesRegistrationAotContribution.java @@ -15,35 +15,72 @@ */ package org.springframework.data.aot; +import java.lang.reflect.Executable; +import java.lang.reflect.Method; import java.util.List; import java.util.function.BiConsumer; +import javax.lang.model.element.Modifier; + +import org.springframework.aot.generate.AccessVisibility; +import org.springframework.aot.generate.GeneratedMethod; import org.springframework.aot.generate.GenerationContext; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; import org.springframework.beans.factory.aot.BeanRegistrationCode; +import org.springframework.beans.factory.aot.BeanRegistrationCodeFragments; +import org.springframework.beans.factory.support.InstanceSupplier; +import org.springframework.beans.factory.support.RegisteredBean; import org.springframework.core.ResolvableType; import org.springframework.data.domain.ManagedTypes; +import org.springframework.data.util.Lazy; +import org.springframework.javapoet.CodeBlock; +import org.springframework.javapoet.MethodSpec.Builder; +import org.springframework.javapoet.ParameterizedTypeName; import org.springframework.lang.Nullable; +import org.springframework.util.ClassUtils; +import org.springframework.util.ObjectUtils; +import org.springframework.util.ReflectionUtils; /** * {@link BeanRegistrationAotContribution} used to contribute a {@link ManagedTypes} registration. + *

+ * Will try to resolve bean definition arguments if possible and fall back to resolving the bean from the context if + * that is not possible. To avoid duplicate invocations of potential scan operations hidden by the {@link ManagedTypes} + * instance the {@link BeanRegistrationAotContribution} will write custom instantiation code via + * {@link BeanRegistrationAotContribution#customizeBeanRegistrationCodeFragments(GenerationContext, BeanRegistrationCodeFragments)}. + * The generated code resolves potential factory methods accepting either a {@link ManagedTypes} instance, or a + * {@link List} of either {@link Class} or {@link String} (classname) values. + * + *

+ * 
+ * public static InstanceSupplier<ManagedTypes> instance() {
+ *   return (registeredBean) -> {
+ *     var types = List.of("com.example.A", "com.example.B");
+ *     return ManagedTypes.ofStream(types.stream().map(it -> ClassUtils.forName(it, registeredBean.getBeanFactory().getBeanClassLoader())));
+ *   }
+ * }
+ * 
+ * 
* * @author John Blum + * @author Christoph Strobl * @see org.springframework.beans.factory.aot.BeanRegistrationAotContribution * @since 3.0.0 */ -public class ManagedTypesRegistrationAotContribution implements BeanRegistrationAotContribution { +public class ManagedTypesRegistrationAotContribution implements RegisteredBeanAotContribution { private final AotContext aotContext; private final ManagedTypes managedTypes; private final BiConsumer contributionAction; + private final RegisteredBean source; public ManagedTypesRegistrationAotContribution(AotContext aotContext, @Nullable ManagedTypes managedTypes, - BiConsumer contributionAction) { + RegisteredBean registeredBean, BiConsumer contributionAction) { this.aotContext = aotContext; this.managedTypes = managedTypes; this.contributionAction = contributionAction; + this.source = registeredBean; } protected AotContext getAotContext() { @@ -63,4 +100,129 @@ public class ManagedTypesRegistrationAotContribution implements BeanRegistration TypeCollector.inspect(types).forEach(type -> contributionAction.accept(type, generationContext)); } } + + @Override + public BeanRegistrationCodeFragments customizeBeanRegistrationCodeFragments(GenerationContext generationContext, + BeanRegistrationCodeFragments codeFragments) { + + if (managedTypes == null) { + return codeFragments; + } + + ManagedTypesInstanceCodeFragment fragment = new ManagedTypesInstanceCodeFragment(getManagedTypes(), source, + codeFragments); + return fragment.canGenerateCode() ? fragment : codeFragments; + } + + @Override + public RegisteredBean getSource() { + return source; + } + + static class ManagedTypesInstanceCodeFragment extends BeanRegistrationCodeFragments { + + private ManagedTypes sourceTypes; + private RegisteredBean source; + private Lazy instanceMethod = Lazy.of(this::findInstanceFactory); + + protected ManagedTypesInstanceCodeFragment(ManagedTypes managedTypes, RegisteredBean source, + BeanRegistrationCodeFragments codeFragments) { + + super(codeFragments); + + this.sourceTypes = managedTypes; + this.source = source; + } + + /** + * @return {@literal true} if the instance method code can be generated. {@literal false} otherwise. + */ + boolean canGenerateCode() { + + if (ObjectUtils.nullSafeEquals(source.getBeanClass(), ManagedTypes.class)) { + return true; + } + return instanceMethod.getNullable() != null; + } + + @Override + public CodeBlock generateInstanceSupplierCode(GenerationContext generationContext, + BeanRegistrationCode beanRegistrationCode, Executable constructorOrFactoryMethod, + boolean allowDirectSupplierShortcut) { + + GeneratedMethod generatedMethod = beanRegistrationCode.getMethods().add("Instance", + this::generateInstanceFactory); + + return CodeBlock.of("$T.$L()", beanRegistrationCode.getClassName(), generatedMethod.getName()); + } + + private CodeBlock toCodeBlock(List> values, boolean allPublic) { + + if (allPublic) { + return CodeBlock.join(values.stream().map(value -> CodeBlock.of("$T.class", value)).toList(), ", "); + } + return CodeBlock.join(values.stream().map(value -> CodeBlock.of("$S", value.getName())).toList(), ", "); + } + + private Method findInstanceFactory() { + + for (Method beanMethod : ReflectionUtils.getDeclaredMethods(source.getBeanClass())) { + + if (beanMethod.getParameterCount() == 1 && java.lang.reflect.Modifier.isPublic(beanMethod.getModifiers()) + && java.lang.reflect.Modifier.isStatic(beanMethod.getModifiers())) { + ResolvableType parameterType = ResolvableType.forMethodParameter(beanMethod, 0, source.getBeanClass()); + if (parameterType.isAssignableFrom(ResolvableType.forType(List.class)) + || parameterType.isAssignableFrom(ResolvableType.forType(ManagedTypes.class))) { + return beanMethod; + } + } + } + return null; + } + + void generateInstanceFactory(Builder method) { + + List> sourceTypes = this.sourceTypes.toList(); + boolean allSourceTypesVisible = sourceTypes.stream() + .allMatch(it -> AccessVisibility.PUBLIC.equals(AccessVisibility.forClass(it))); + + ParameterizedTypeName targetTypeName = ParameterizedTypeName.get(InstanceSupplier.class, source.getBeanClass()); + + method.addModifiers(Modifier.PRIVATE, Modifier.STATIC); + method.returns(targetTypeName); + + CodeBlock.Builder builder = CodeBlock.builder().add("return ").beginControlFlow("(registeredBean -> "); + + builder.addStatement("var types = $T.of($L)", List.class, toCodeBlock(sourceTypes, allSourceTypesVisible)); + + if (allSourceTypesVisible) { + builder.addStatement("var managedTypes = $T.fromIterable($L)", ManagedTypes.class, "types"); + } else { + builder.add(CodeBlock.builder() + .beginControlFlow("var managedTypes = $T.fromStream(types.stream().map(it ->", ManagedTypes.class) + .beginControlFlow("try") + .addStatement("return $T.forName(it, registeredBean.getBeanFactory().getBeanClassLoader())", + ClassUtils.class) + .nextControlFlow("catch ($T e)", ClassNotFoundException.class) + .addStatement("throw new $T($S, e)", IllegalArgumentException.class, "Cannot to load type").endControlFlow() + .endControlFlow("))").build()); + } + if (ObjectUtils.nullSafeEquals(source.getBeanClass(), ManagedTypes.class)) { + builder.add("return managedTypes"); + } else { + Method instanceFactoryMethod = instanceMethod.get(); + if (ResolvableType.forMethodParameter(instanceFactoryMethod, 0) + .isAssignableFrom(ResolvableType.forType(ManagedTypes.class))) { + builder.addStatement("return $T.$L($L)", instanceFactoryMethod.getDeclaringClass(), + instanceFactoryMethod.getName(), "managedTypes"); + + } else { + builder.addStatement("return $T.$L($L.toList())", instanceFactoryMethod.getDeclaringClass(), + instanceFactoryMethod.getName(), "managedTypes"); + } + } + builder.endControlFlow(")"); + method.addCode(builder.build()); + } + } } diff --git a/src/main/java/org/springframework/data/aot/RegisteredBeanAotContribution.java b/src/main/java/org/springframework/data/aot/RegisteredBeanAotContribution.java new file mode 100644 index 000000000..b35567546 --- /dev/null +++ b/src/main/java/org/springframework/data/aot/RegisteredBeanAotContribution.java @@ -0,0 +1,27 @@ +/* + * Copyright 2022 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.aot; + +import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; +import org.springframework.beans.factory.support.RegisteredBean; + +/** + * @author Christoph Strobl + */ +public interface RegisteredBeanAotContribution extends BeanRegistrationAotContribution { + + RegisteredBean getSource(); +} diff --git a/src/test/java/org/springframework/data/aot/AotTestCodeContributionBuilder.java b/src/test/java/org/springframework/data/aot/AotTestCodeContributionBuilder.java new file mode 100644 index 000000000..b924209cf --- /dev/null +++ b/src/test/java/org/springframework/data/aot/AotTestCodeContributionBuilder.java @@ -0,0 +1,92 @@ +/* + * Copyright 2022 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.aot; + +import java.util.function.Consumer; + +import javax.lang.model.element.Modifier; + +import org.mockito.Mockito; +import org.springframework.aot.test.generate.TestGenerationContext; +import org.springframework.aot.test.generate.compile.Compiled; +import org.springframework.aot.test.generate.compile.TestCompiler; +import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; +import org.springframework.beans.factory.aot.BeanRegistrationCodeFragments; +import org.springframework.beans.factory.support.InstanceSupplier; +import org.springframework.javapoet.CodeBlock; +import org.springframework.javapoet.MethodSpec; +import org.springframework.javapoet.ParameterizedTypeName; + +/** + * @author Christoph Strobl + */ +public class AotTestCodeContributionBuilder { + + TestGenerationContext generationContext; + MockBeanRegistrationCode beanRegistrationCode; + BeanRegistrationAotContribution contribution; + + static AotTestCodeContributionBuilder withContextFor(Class type) { + return withContext(new TestGenerationContext(type)); + } + + static AotTestCodeContributionBuilder withContext(TestGenerationContext ctx) { + + AotTestCodeContributionBuilder codeGenerationBuilder = new AotTestCodeContributionBuilder(); + codeGenerationBuilder.generationContext = ctx; + codeGenerationBuilder.beanRegistrationCode = new MockBeanRegistrationCode(ctx); + return codeGenerationBuilder; + } + + BeanRegistrationCodeFragments getFragments(BeanRegistrationAotContribution contribution) { + + this.contribution = contribution; + + return contribution.customizeBeanRegistrationCodeFragments(generationContext, + Mockito.mock(BeanRegistrationCodeFragments.class)); + } + + AotTestCodeContributionBuilder writeContentFor(BeanRegistrationAotContribution contribution) { + + CodeBlock codeBlock = getFragments(contribution).generateInstanceSupplierCode(generationContext, + beanRegistrationCode, null, false); + + Class beanType = Object.class; + try { + beanType = contribution instanceof RegisteredBeanAotContribution + ? ((RegisteredBeanAotContribution) contribution).getSource().getBeanClass() + : Object.class; + } catch (Exception e) {} + + ParameterizedTypeName parameterizedReturnTypeName = ParameterizedTypeName.get(InstanceSupplier.class, beanType); + beanRegistrationCode.getTypeBuilder().set(type -> { + type.addModifiers(Modifier.PUBLIC); + type.addMethod(MethodSpec.methodBuilder("get").addModifiers(Modifier.PUBLIC).returns(parameterizedReturnTypeName) + .addStatement("return $L", codeBlock).build()); + }); + + return this; + } + + public void compile() { + compile(it -> {}); + } + + public void compile(Consumer compiled) { + generationContext.writeGeneratedContent(); + TestCompiler.forSystem().withFiles(generationContext.getGeneratedFiles()).compile(compiled); + } +} diff --git a/src/test/java/org/springframework/data/aot/BeanRegistrationContributionAssert.java b/src/test/java/org/springframework/data/aot/BeanRegistrationContributionAssert.java index d500c072d..8564c3588 100644 --- a/src/test/java/org/springframework/data/aot/BeanRegistrationContributionAssert.java +++ b/src/test/java/org/springframework/data/aot/BeanRegistrationContributionAssert.java @@ -20,7 +20,6 @@ import static org.mockito.Mockito.*; import java.util.function.Consumer; import org.assertj.core.api.AbstractAssert; - import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.test.generate.TestGenerationContext; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; diff --git a/src/test/java/org/springframework/data/aot/DeferredTypeBuilder.java b/src/test/java/org/springframework/data/aot/DeferredTypeBuilder.java new file mode 100644 index 000000000..997d86a3f --- /dev/null +++ b/src/test/java/org/springframework/data/aot/DeferredTypeBuilder.java @@ -0,0 +1,42 @@ +/* + * Copyright 2022 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.aot; + +import java.util.function.Consumer; + +import org.springframework.javapoet.TypeSpec; +import org.springframework.javapoet.TypeSpec.Builder; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; + +/** + * @author Christoph Strobl + */ +public class DeferredTypeBuilder implements Consumer { + + @Nullable + private Consumer type; + + @Override + public void accept(Builder type) { + Assert.notNull(this.type, "No type builder set"); + this.type.accept(type); + } + + public void set(Consumer type) { + this.type = type; + } +} diff --git a/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java b/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java index 56d083110..4fa54d870 100644 --- a/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java +++ b/src/test/java/org/springframework/data/aot/ManagedTypesBeanRegistrationAotProcessorUnitTests.java @@ -19,21 +19,32 @@ import static org.assertj.core.api.Assertions.*; import static org.mockito.ArgumentMatchers.*; import static org.mockito.Mockito.*; +import java.util.Arrays; import java.util.Collections; +import java.util.concurrent.atomic.AtomicInteger; import java.util.function.Consumer; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.mockito.Mockito; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.hint.predicate.RuntimeHintsPredicates; import org.springframework.aot.test.generate.TestGenerationContext; import org.springframework.beans.factory.BeanCreationException; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; +import org.springframework.beans.factory.aot.BeanRegistrationCodeFragments; import org.springframework.beans.factory.support.BeanDefinitionBuilder; import org.springframework.beans.factory.support.DefaultListableBeanFactory; +import org.springframework.beans.factory.support.InstanceSupplier; import org.springframework.beans.factory.support.RegisteredBean; import org.springframework.beans.factory.support.RootBeanDefinition; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.data.aot.ManagedTypesRegistrationAotContribution.ManagedTypesInstanceCodeFragment; import org.springframework.data.domain.ManagedTypes; +import org.springframework.javapoet.MethodSpec; +import org.springframework.javapoet.MethodSpec.Builder; +import org.springframework.test.util.ReflectionTestUtils; /** * @author Christoph Strobl @@ -47,6 +58,9 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { final RootBeanDefinition myManagedTypesDefinition = (RootBeanDefinition) BeanDefinitionBuilder .rootBeanDefinition(MyManagedTypes.class).getBeanDefinition(); + final RootBeanDefinition invocationCountingManagedTypesDefinition = (RootBeanDefinition) BeanDefinitionBuilder + .rootBeanDefinition(InvocationRecordingManagedTypes.class).getBeanDefinition(); + DefaultListableBeanFactory beanFactory; @BeforeEach @@ -154,6 +168,79 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { verify(beanFactory).getBean(eq("commons.managed-types"), eq(ManagedTypes.class)); } + @Test // GH-2680 + void generatesInstanceSupplierCodeFragmentToAvoidDuplicateInvocations() { + + beanFactory.registerBeanDefinition("commons.managed-types", invocationCountingManagedTypesDefinition); + RegisteredBean registeredBean = RegisteredBean.of(beanFactory, "commons.managed-types"); + + BeanRegistrationAotContribution contribution = createPostProcessor("commons") + .processAheadOfTime(RegisteredBean.of(beanFactory, "commons.managed-types")); + + AotTestCodeContributionBuilder.withContextFor(this.getClass()).writeContentFor(contribution).compile(it -> { + + InvocationRecordingManagedTypes sourceTypes = beanFactory.getBean(InvocationRecordingManagedTypes.class); + assertThat(sourceTypes.getCounter()).isOne(); + + InstanceSupplier types = ReflectionTestUtils + .invokeMethod(it.getAllCompiledClasses().iterator().next(), "instance"); + try { + assertThat(types.get(registeredBean).source).isNotSameAs(sourceTypes); + } catch (Exception e) { + throw new RuntimeException(e); + } + }); + } + + @Test // GH-2680 + void generatesInstanceSupplierCodeFragmentForTypeWithCustomFactoryMethod() { + + beanFactory.registerBeanDefinition("commons.managed-types", + BeanDefinitionBuilder.rootBeanDefinition(StoreManagedTypesWithCustomFactoryMethod.class).getBeanDefinition()); + + RegisteredBean registeredBean = RegisteredBean.of(beanFactory, "commons.managed-types"); + + BeanRegistrationAotContribution contribution = createPostProcessor("commons").processAheadOfTime(registeredBean); + + AotTestCodeContributionBuilder.withContextFor(this.getClass()).writeContentFor(contribution).compile(it -> { + + InstanceSupplier types = ReflectionTestUtils + .invokeMethod(it.getAllCompiledClasses().iterator().next(), "instance"); + + try { + assertThat(types.get(registeredBean).toList()).containsExactlyInAnyOrder(A.class, B.class); + } catch (Exception e) { + throw new RuntimeException(e); + } + }); + } + + @Test // GH-2680 + void canGenerateCodeReturnsTrueIfFactoryMethodPresent() { + + beanFactory.registerBeanDefinition("managed-types", managedTypesDefinition); + RegisteredBean registeredBean = RegisteredBean.of(beanFactory, "managed-types"); + + ManagedTypesInstanceCodeFragment fragment = new ManagedTypesInstanceCodeFragment( + ManagedTypes.from(A.class, B.class), registeredBean, Mockito.mock(BeanRegistrationCodeFragments.class)); + Builder methodBuilder = MethodSpec.methodBuilder("instance"); + fragment.generateInstanceFactory(methodBuilder); + + assertThat(fragment.canGenerateCode()).isTrue(); + } + + @Test // GH-2680 + void canGenerateCodeReturnsFalseIfNoFactoryMethodPresent() { + + beanFactory.registerBeanDefinition("managed-types", myManagedTypesDefinition); + RegisteredBean registeredBean = RegisteredBean.of(beanFactory, "managed-types"); + + ManagedTypesInstanceCodeFragment fragment = new ManagedTypesInstanceCodeFragment( + ManagedTypes.from(A.class, B.class), registeredBean, Mockito.mock(BeanRegistrationCodeFragments.class)); + + assertThat(fragment.canGenerateCode()).isFalse(); + } + private ManagedTypesBeanRegistrationAotProcessor createPostProcessor(String moduleIdentifier) { ManagedTypesBeanRegistrationAotProcessor postProcessor = new ManagedTypesBeanRegistrationAotProcessor(); postProcessor.setModuleIdentifier(moduleIdentifier); @@ -172,5 +259,67 @@ class ManagedTypesBeanRegistrationAotProcessorUnitTests { } } + static class StoreManagedTypesWithFactoryMethodOfClassNames implements ManagedTypes { + @Override + public void forEach(Consumer> action) { + // just do nothing ¯\_(ツ)_/¯ + } + } + + public static class StoreManagedTypesWithCustomFactoryMethod implements ManagedTypes { + + private ManagedTypes source; + + public StoreManagedTypesWithCustomFactoryMethod() { + source = it -> Arrays.asList(A.class, B.class).forEach(it); + } + + public StoreManagedTypesWithCustomFactoryMethod(ManagedTypes source) { + this.source = source; + } + + public static StoreManagedTypesWithCustomFactoryMethod of(ManagedTypes source) { + return new StoreManagedTypesWithCustomFactoryMethod(source); + } + + @Override + public void forEach(Consumer> action) { + source.forEach(action); + } + } + + public static class InvocationRecordingManagedTypes implements ManagedTypes { + + private AtomicInteger counter = new AtomicInteger(0); + private ManagedTypes source = ManagedTypes.from(A.class, B.class); + + public static InvocationRecordingManagedTypes from(ManagedTypes source) { + + InvocationRecordingManagedTypes newInstance = new InvocationRecordingManagedTypes(); + newInstance.source = source; + return newInstance; + } + + @Override + public void forEach(Consumer> action) { + + counter.getAndIncrement(); + source.forEach(action); + } + + public int getCounter() { + return counter.get(); + } + } + static class NotManagedTypes {} + + @Configuration(proxyBeanMethods = false) + public static class EntityManagerWithPackagesToScanConfiguration { + + @Bean(name = "commons.managed-types") + ManagedTypes managedTypes() { + return ManagedTypes.from(A.class, B.class); + } + } } diff --git a/src/test/java/org/springframework/data/aot/MockBeanRegistrationCode.java b/src/test/java/org/springframework/data/aot/MockBeanRegistrationCode.java new file mode 100644 index 000000000..5b909e92c --- /dev/null +++ b/src/test/java/org/springframework/data/aot/MockBeanRegistrationCode.java @@ -0,0 +1,61 @@ +/* + * Copyright 2022 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.aot; + +import java.util.ArrayList; +import java.util.List; + +import org.springframework.aot.generate.GeneratedClass; +import org.springframework.aot.generate.GeneratedMethods; +import org.springframework.aot.generate.GenerationContext; +import org.springframework.aot.generate.MethodReference; +import org.springframework.beans.factory.aot.BeanRegistrationCode; +import org.springframework.javapoet.ClassName; + +/** + * @author Christoph Strobl + */ +public class MockBeanRegistrationCode implements BeanRegistrationCode { + + private final GeneratedClass generatedClass; + + private final List instancePostProcessors = new ArrayList<>(); + + private final DeferredTypeBuilder typeBuilder = new DeferredTypeBuilder(); + + public MockBeanRegistrationCode(GenerationContext generationContext) { + this.generatedClass = generationContext.getGeneratedClasses().addForFeature("TestCode", this.typeBuilder); + } + + public DeferredTypeBuilder getTypeBuilder() { + return this.typeBuilder; + } + + @Override + public ClassName getClassName() { + return this.generatedClass.getName(); + } + + @Override + public GeneratedMethods getMethods() { + return this.generatedClass.getMethods(); + } + + @Override + public void addInstancePostProcessor(MethodReference methodReference) { + this.instancePostProcessors.add(methodReference); + } +}