From 281f18641144d454d7466b587a9138ae43c428f5 Mon Sep 17 00:00:00 2001 From: Christoph Strobl Date: Fri, 20 Sep 2024 11:36:54 +0200 Subject: [PATCH] Add support for AOT generated repository implementations and wire fragments in BeanDefinition AOT code. We now provide infrastructure to generate AOT repository method code that implements Query method behavior. No longer use spring.factories but write some custom bean config code so that one of the properties can provide an instance of the generated repository. Closes #3265 --- .../springframework/data/aot/AotContext.java | 7 + .../aot/generate/AotCodeContributor.java | 25 +++ .../aot/generate/AotRepositoryBuilder.java | 169 ++++++++++++++ .../AotRepositoryConstructorBuilder.java | 94 ++++++++ .../AotRepositoryImplementationMetadata.java | 91 ++++++++ .../generate/AotRepositoryMethodBuilder.java | 140 ++++++++++++ .../AotRepositoryMethodGenerationContext.java | 207 ++++++++++++++++++ ...epositoryMethodImplementationMetadata.java | 74 +++++++ .../repository/aot/generate/CodeBlocks.java | 65 ++++++ .../aot/generate/RepositoryContributor.java | 92 ++++++++ .../config/AotRepositoryContext.java | 2 +- .../config/AotRepositoryInformation.java | 9 +- ...RepositoryRegistrationAotContribution.java | 70 +++++- .../RepositoryRegistrationAotProcessor.java | 18 +- .../support/DefaultRepositoryInformation.java | 21 ++ .../support/RepositoryFactoryBeanSupport.java | 16 ++ .../support/RepositoryFactorySupport.java | 10 + src/test/java/example/UserRepository.java | 46 ++++ .../springframework/data/DependencyTests.java | 2 + .../data/aot/CodeContributionAssert.java | 13 +- .../DummyModuleAotRepositoryContext.java | 105 +++++++++ ...ModuleDefaultRepositoryImplementation.java | 28 +++ .../generate/RepositoryBuilderUnitTests.java | 37 ++++ .../RepositoryContributorUnitTests.java | 62 ++++++ .../generate/StubRepositoryInformation.java | 127 +++++++++++ ...DefaultRepositoryInformationUnitTests.java | 16 ++ 26 files changed, 1535 insertions(+), 11 deletions(-) create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotCodeContributor.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryImplementationMetadata.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodBuilder.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodGenerationContext.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodImplementationMetadata.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/CodeBlocks.java create mode 100644 src/main/java/org/springframework/data/repository/aot/generate/RepositoryContributor.java create mode 100644 src/test/java/example/UserRepository.java create mode 100644 src/test/java/org/springframework/data/repository/aot/generate/DummyModuleAotRepositoryContext.java create mode 100644 src/test/java/org/springframework/data/repository/aot/generate/DummyModuleDefaultRepositoryImplementation.java create mode 100644 src/test/java/org/springframework/data/repository/aot/generate/RepositoryBuilderUnitTests.java create mode 100644 src/test/java/org/springframework/data/repository/aot/generate/RepositoryContributorUnitTests.java create mode 100644 src/test/java/org/springframework/data/repository/aot/generate/StubRepositoryInformation.java diff --git a/src/main/java/org/springframework/data/aot/AotContext.java b/src/main/java/org/springframework/data/aot/AotContext.java index 4d8c0baed..6bb6f2a56 100644 --- a/src/main/java/org/springframework/data/aot/AotContext.java +++ b/src/main/java/org/springframework/data/aot/AotContext.java @@ -30,6 +30,7 @@ import org.springframework.beans.factory.config.BeanDefinition; import org.springframework.beans.factory.config.BeanReference; import org.springframework.beans.factory.config.ConfigurableListableBeanFactory; import org.springframework.beans.factory.support.RootBeanDefinition; +import org.springframework.core.SpringProperties; import org.springframework.data.util.TypeScanner; import org.springframework.util.Assert; @@ -49,6 +50,12 @@ import org.springframework.util.Assert; */ public interface AotContext { + String GENERATED_REPOSITORIES_ENABLED = "spring.aot.repositories.enabled"; + + static boolean aotGeneratedRepositoriesEnabled() { + return SpringProperties.getFlag(GENERATED_REPOSITORIES_ENABLED); + } + /** * Create an {@link AotContext} backed by the given {@link BeanFactory}. * diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotCodeContributor.java b/src/main/java/org/springframework/data/repository/aot/generate/AotCodeContributor.java new file mode 100644 index 000000000..350afa768 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotCodeContributor.java @@ -0,0 +1,25 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import org.springframework.aot.generate.GenerationContext; + +/** + * @author Christoph Strobl + */ +public interface AotCodeContributor { + void contribute(GenerationContext generationContext); +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java new file mode 100644 index 000000000..7cee91e2e --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java @@ -0,0 +1,169 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import java.time.YearMonth; +import java.time.ZoneId; +import java.time.temporal.ChronoField; +import java.util.Map; +import java.util.function.Consumer; +import java.util.function.Function; + +import javax.lang.model.element.Modifier; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.aot.generate.ClassNameGenerator; +import org.springframework.aot.generate.Generated; +import org.springframework.data.repository.CrudRepository; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.javapoet.ClassName; +import org.springframework.javapoet.FieldSpec; +import org.springframework.javapoet.JavaFile; +import org.springframework.javapoet.TypeName; +import org.springframework.javapoet.TypeSpec; +import org.springframework.stereotype.Component; +import org.springframework.util.ReflectionUtils; + +/** + * @author Christoph Strobl + */ +public class AotRepositoryBuilder { + + private final RepositoryInformation repositoryInformation; + private final AotRepositoryImplementationMetadata generationMetadata; + + private Consumer constructorBuilderCustomizer; + private Function methodContextFunction; + private RepositoryCustomizer customizer; + + public static AotRepositoryBuilder forRepository(RepositoryInformation repositoryInformation) { + return new AotRepositoryBuilder(repositoryInformation); + } + + AotRepositoryBuilder(RepositoryInformation repositoryInformation) { + + this.repositoryInformation = repositoryInformation; + this.generationMetadata = new AotRepositoryImplementationMetadata(className()); + this.generationMetadata.addField(FieldSpec + .builder(TypeName.get(Log.class), "logger", Modifier.PRIVATE, Modifier.STATIC, Modifier.FINAL) + .initializer("$T.getLog($T.class)", TypeName.get(LogFactory.class), this.generationMetadata.getTargetTypeName()) + .build()); + + this.customizer = (info, metadata, builder) -> {}; + } + + public JavaFile javaFile() { + + YearMonth creationDate = YearMonth.now(ZoneId.of("UTC")); + + // start creating the type + TypeSpec.Builder builder = TypeSpec.classBuilder(this.generationMetadata.getTargetTypeName()) // + .addModifiers(Modifier.PUBLIC) // + .addAnnotation(Generated.class) // + .addJavadoc("AOT generated repository implementation for {@link $T}.\n", + repositoryInformation.getRepositoryInterface()) // + .addJavadoc("\n") // + .addJavadoc("@since $L/$L\n", creationDate.get(ChronoField.YEAR), creationDate.get(ChronoField.MONTH_OF_YEAR)) // + .addJavadoc("@author $L", "Spring Data"); // TODO: does System.getProperty("user.name") make sense here? + + // TODO: we do not need that here + // .addSuperinterface(repositoryInformation.getRepositoryInterface()); + + // create the constructor + AotRepositoryConstructorBuilder constructorBuilder = new AotRepositoryConstructorBuilder(repositoryInformation, + generationMetadata); + constructorBuilderCustomizer.accept(constructorBuilder); + builder.addMethod(constructorBuilder.buildConstructor()); + + // write methods + // start with the derived ones + ReflectionUtils.doWithMethods(repositoryInformation.getRepositoryInterface(), method -> { + + AotRepositoryMethodGenerationContext context = new AotRepositoryMethodGenerationContext(method, + repositoryInformation, generationMetadata); + AotRepositoryMethodBuilder methodBuilder = methodContextFunction.apply(context); + if (methodBuilder != null) { + builder.addMethod(methodBuilder.buildMethod()); + } + + }, it -> { + + /* + the isBaseClassMethod(it) check seems to have some issues. + need to hard code it here + */ + + if (ReflectionUtils.findMethod(CrudRepository.class, it.getName(), it.getParameterTypes()) != null) { + return false; + } + + return !repositoryInformation.isBaseClassMethod(it) && !repositoryInformation.isCustomMethod(it) + && !it.isDefault(); + }); + + // write fields at the end so we make sure to capture things added by methods + generationMetadata.getFields().values().forEach(builder::addField); + + // finally customize the file itself + this.customizer.customize(repositoryInformation, generationMetadata, builder); + return JavaFile.builder(packageName(), builder.build()).build(); + } + + AotRepositoryBuilder withConstructorCustomizer(Consumer constuctorBuilder) { + + this.constructorBuilderCustomizer = constuctorBuilder; + return this; + } + + AotRepositoryBuilder withDerivedMethodFunction( + Function methodContextFunction) { + this.methodContextFunction = methodContextFunction; + return this; + } + + AotRepositoryBuilder withFileCustomizer(RepositoryCustomizer repositoryCustomizer) { + + this.customizer = repositoryCustomizer; + return this; + } + + AotRepositoryImplementationMetadata getGenerationMetadata() { + return generationMetadata; + } + + private ClassName className() { + return new ClassNameGenerator(ClassName.get(packageName(), typeName())).generateClassName("Aot", null); + } + + private String packageName() { + return repositoryInformation.getRepositoryInterface().getPackageName(); + } + + private String typeName() { + return "%sImpl".formatted(repositoryInformation.getRepositoryInterface().getSimpleName()); + } + + Map getAutowireFields() { + return generationMetadata.getConstructorArguments(); + } + + public interface RepositoryCustomizer { + + void customize(RepositoryInformation repositoryInformation, AotRepositoryImplementationMetadata metadata, + TypeSpec.Builder builder); + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java new file mode 100644 index 000000000..c1d592150 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java @@ -0,0 +1,94 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import java.util.List; +import java.util.Map.Entry; + +import javax.lang.model.element.Modifier; + +import org.springframework.core.ResolvableType; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.javapoet.MethodSpec; +import org.springframework.javapoet.ParameterizedTypeName; +import org.springframework.javapoet.TypeName; + +/** + * @author Christoph Strobl + */ +public class AotRepositoryConstructorBuilder { + + private final RepositoryInformation repositoryInformation; + private final AotRepositoryImplementationMetadata metadata; + + private ConstructorCustomizer customizer = (info, builder) -> {}; + + AotRepositoryConstructorBuilder(RepositoryInformation repositoryInformation, + AotRepositoryImplementationMetadata metadata) { + + this.repositoryInformation = repositoryInformation; + this.metadata = metadata; + } + + public void addParameter(String parameterName, Class type) { + + ResolvableType resolvableType = ResolvableType.forClass(type); + if (!resolvableType.hasGenerics() || !resolvableType.hasResolvableGenerics()) { + addParameter(parameterName, TypeName.get(type)); + return; + } + addParameter(parameterName, ParameterizedTypeName.get(type, resolvableType.resolveGenerics())); + } + + public void addParameter(String parameterName, TypeName type) { + + this.metadata.addConstructorArgument(parameterName, type); + this.metadata.addField(parameterName, type, Modifier.PRIVATE, Modifier.FINAL); + } + + public void customize(ConstructorCustomizer customizer) { + this.customizer = customizer; + } + + MethodSpec buildConstructor() { + + MethodSpec.Builder builder = MethodSpec.constructorBuilder().addModifiers(Modifier.PUBLIC); + for (Entry parameter : this.metadata.getConstructorArguments().entrySet()) { + builder.addParameter(parameter.getValue(), parameter.getKey()).addStatement("this.$N = $N", parameter.getKey(), + parameter.getKey()); + } + customizer.customize(repositoryInformation, builder); + return builder.build(); + } + + private static TypeName getDefaultStoreRepositoryImplementationType(RepositoryInformation repositoryInformation) { + + ResolvableType resolvableType = ResolvableType.forClass(repositoryInformation.getRepositoryBaseClass()); + if (resolvableType.hasGenerics()) { + List> generics = List.of(); + if (resolvableType.getGenerics().length == 2) { // TODO: Find some other way to resolve generics + generics = List.of(repositoryInformation.getDomainType(), repositoryInformation.getIdType()); + } + return ParameterizedTypeName.get(repositoryInformation.getRepositoryBaseClass(), generics.toArray(Class[]::new)); + } + return TypeName.get(repositoryInformation.getRepositoryBaseClass()); + } + + public interface ConstructorCustomizer { + + void customize(RepositoryInformation repositoryInformation, MethodSpec.Builder builder); + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryImplementationMetadata.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryImplementationMetadata.java new file mode 100644 index 000000000..c9b004c08 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryImplementationMetadata.java @@ -0,0 +1,91 @@ +/* + * Copyright 2025 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.repository.aot.generate; + +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Map.Entry; + +import javax.lang.model.element.Modifier; + +import org.springframework.javapoet.ClassName; +import org.springframework.javapoet.FieldSpec; +import org.springframework.javapoet.TypeName; +import org.springframework.lang.Nullable; + +/** + * @author Christoph Strobl + */ +class AotRepositoryImplementationMetadata { + + private ClassName className; + private Map fields = new HashMap<>(3); + private final Map constructorArguments = new LinkedHashMap<>(3); + + public AotRepositoryImplementationMetadata(ClassName className) { + this.className = className; + } + + @Nullable + public String fieldNameOf(Class type) { + + TypeName lookup = TypeName.get(type).withoutAnnotations(); + for (Entry field : fields.entrySet()) { + if (field.getValue().type.withoutAnnotations().equals(lookup)) { + return field.getKey(); + } + } + + return null; + } + + public ClassName getTargetTypeName() { + return className; + } + + public String getTargetTypeSimpleName() { + return className.simpleName(); + } + + public String getTargetTypePackageName() { + return className.packageName(); + } + + public boolean hasField(String fieldName) { + return fields.containsKey(fieldName); + } + + public void addField(String fieldName, TypeName type, Modifier... modifiers) { + fields.put(fieldName, FieldSpec.builder(type, fieldName, modifiers).build()); + } + + public void addField(FieldSpec fieldSpec) { + fields.put(fieldSpec.name, fieldSpec); + } + + Map getFields() { + return fields; + } + + public Map getConstructorArguments() { + return constructorArguments; + } + + public void addConstructorArgument(String parameterName, TypeName type) { + this.constructorArguments.put(parameterName, type); + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodBuilder.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodBuilder.java new file mode 100644 index 000000000..f4ca54264 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodBuilder.java @@ -0,0 +1,140 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import java.lang.reflect.Method; +import java.lang.reflect.Parameter; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Map.Entry; +import java.util.stream.Collectors; + +import javax.lang.model.element.Modifier; + +import org.springframework.core.MethodParameter; +import org.springframework.core.ResolvableType; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.javapoet.MethodSpec; +import org.springframework.javapoet.ParameterSpec; +import org.springframework.javapoet.ParameterizedTypeName; +import org.springframework.javapoet.TypeName; +import org.springframework.lang.Nullable; +import org.springframework.util.StringUtils; + +/** + * @author Christoph Strobl + */ +public class AotRepositoryMethodBuilder { + + private final AotRepositoryMethodGenerationContext context; + + private RepositoryMethodCustomizer customizer = (context, body) -> {}; + + public AotRepositoryMethodBuilder(AotRepositoryMethodGenerationContext context) { + + this.context = context; + initReturnType(context.getMethod(), context.getRepositoryInformation()); + initParameters(context.getMethod(), context.getRepositoryInformation()); + } + + public void addParameter(String parameterName, Class type) { + + ResolvableType resolvableType = ResolvableType.forClass(type); + if (!resolvableType.hasGenerics() || !resolvableType.hasResolvableGenerics()) { + addParameter(parameterName, TypeName.get(type)); + return; + } + addParameter(parameterName, ParameterizedTypeName.get(type, resolvableType.resolveGenerics())); + } + + public void addParameter(String parameterName, TypeName type) { + addParameter(ParameterSpec.builder(type, parameterName).build()); + } + + public void addParameter(ParameterSpec parameter) { + this.context.addParameter(parameter); + } + + public void setReturnType(@Nullable TypeName returnType, @Nullable TypeName actualReturnType) { + this.context.getTargetMethodMetadata().setReturnType(returnType); + this.context.getTargetMethodMetadata().setActualReturnType(actualReturnType); + } + + public AotRepositoryMethodBuilder customize(RepositoryMethodCustomizer customizer) { + this.customizer = customizer; + return this; + } + + MethodSpec buildMethod() { + + MethodSpec.Builder builder = MethodSpec.methodBuilder(context.getMethod().getName()).addModifiers(Modifier.PUBLIC); + if (!context.returnsVoid()) { + builder.returns(context.getReturnType()); + } + builder.addJavadoc("AOT generated implementation of {@link $T#$L($L)}.", context.getMethod().getDeclaringClass(), + context.getMethod().getName(), + StringUtils.collectionToCommaDelimitedString(context.getTargetMethodMetadata().getMethodArguments().values().stream() + .map(it -> it.type.toString()).collect(Collectors.toList()))); + context.getTargetMethodMetadata().getMethodArguments().forEach((name, spec) -> builder.addParameter(spec)); + customizer.customize(context, builder); + return builder.build(); + } + + private void initParameters(Method method, RepositoryInformation repositoryInformation) { + + ResolvableType repositoryInterface = ResolvableType.forClass(repositoryInformation.getRepositoryInterface()); + if (method.getParameterCount() > 0) { + int index = 0; + for (Parameter parameter : method.getParameters()) { + + ResolvableType resolvableParameterType = ResolvableType.forMethodParameter(new MethodParameter(method, index), + repositoryInterface); + + TypeName parameterType = TypeName.get(resolvableParameterType.resolve()); + if (resolvableParameterType.hasGenerics()) { + parameterType = ParameterizedTypeName.get(resolvableParameterType.resolve(), + resolvableParameterType.resolveGenerics()); + } + addParameter(parameter.getName(), parameterType); + index++; + } + + } + } + + private void initReturnType(Method method, RepositoryInformation repositoryInformation) { + + ResolvableType returnType = ResolvableType.forMethodReturnType(method, + repositoryInformation.getRepositoryInterface()); + + TypeName returnTypeName = TypeName.get(returnType.resolve()); + TypeName actualReturnTypeName = null; + if (returnType.hasGenerics()) { + Class[] generics = returnType.resolveGenerics(); + returnTypeName = ParameterizedTypeName.get(returnType.resolve(), generics); + + if (generics.length == 1) { + actualReturnTypeName = TypeName.get(generics[0]); + } + } + + setReturnType(returnTypeName, actualReturnTypeName); + } + + public interface RepositoryMethodCustomizer { + void customize(AotRepositoryMethodGenerationContext context, MethodSpec.Builder builder); + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodGenerationContext.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodGenerationContext.java new file mode 100644 index 000000000..fe3540171 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodGenerationContext.java @@ -0,0 +1,207 @@ +/* + * Copyright 2025. 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 + * + * http://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. + */ + +/* + * Copyright 2025 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 + * + * http://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.repository.aot.generate; + +import java.lang.annotation.Annotation; +import java.lang.reflect.Method; +import java.util.ArrayList; +import java.util.Collection; +import java.util.Map.Entry; +import java.util.Optional; + +import javax.lang.model.element.Modifier; + +import org.springframework.core.annotation.AnnotatedElementUtils; +import org.springframework.core.annotation.AnnotationAttributes; +import org.springframework.data.domain.Limit; +import org.springframework.data.domain.Page; +import org.springframework.data.domain.Pageable; +import org.springframework.data.domain.Slice; +import org.springframework.data.domain.Sort; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.data.repository.query.parser.PartTree; +import org.springframework.javapoet.FieldSpec; +import org.springframework.javapoet.ParameterSpec; +import org.springframework.javapoet.TypeName; +import org.springframework.lang.Nullable; +import org.springframework.util.ClassUtils; + +/** + * @author Christoph Strobl + * @since 2025/01 + */ +public class AotRepositoryMethodGenerationContext { + + private final Method method; + private final RepositoryInformation repositoryInformation; + private final AotRepositoryImplementationMetadata targetTypeMetadata; + private final AotRepositoryMethodImplementationMetadata targetMethodMetadata; + private final CodeBlocks codeBlocks; + @Nullable PartTree partTree; + + public AotRepositoryMethodGenerationContext(Method method, RepositoryInformation repositoryInformation, + AotRepositoryImplementationMetadata targetTypeMetadata) { + + this.method = method; + this.repositoryInformation = repositoryInformation; + this.targetTypeMetadata = targetTypeMetadata; + this.targetMethodMetadata = new AotRepositoryMethodImplementationMetadata(); + this.codeBlocks = new CodeBlocks(targetTypeMetadata); + try { + this.partTree = new PartTree(method.getName(), repositoryInformation.getDomainType()); + } catch (Exception e) { + // not a part tree quer + } + } + + public boolean hasField(String fieldName) { + return targetTypeMetadata.hasField(fieldName); + } + + public void addField(String fieldName, TypeName type, Modifier... modifiers) { + targetTypeMetadata.addField(fieldName, type, modifiers); + } + + public void addField(FieldSpec fieldSpec) { + targetTypeMetadata.addField(fieldSpec); + } + + public String fieldNameOf(Class type) { + return targetTypeMetadata.fieldNameOf(type); + } + + public RepositoryInformation getRepositoryInformation() { + return repositoryInformation; + } + + public Method getMethod() { + return method; + } + + AotRepositoryImplementationMetadata getTargetTypeMetadata() { + return targetTypeMetadata; + } + + @Nullable + public String getParameterNameOf(Class type) { + return targetMethodMetadata.getParameterNameOf(type); + } + + public String getParameterNameOfPosition(int position) { + + ArrayList> entries = new ArrayList<>( + targetMethodMetadata.getMethodArguments().entrySet()); + if (position < entries.size()) { + return entries.get(position).getKey(); + } + return null; + } + + public void addParameter(ParameterSpec parameter) { + this.targetMethodMetadata.addParameter(parameter); + } + + public boolean returnsVoid() { + return getMethod().getReturnType().equals(Void.TYPE); + } + + public boolean returnsPage() { + return ClassUtils.isAssignable(Page.class, getMethod().getReturnType()); + } + + public boolean returnsSlice() { + return ClassUtils.isAssignable(Slice.class, getMethod().getReturnType()); + } + + public boolean returnsCollection() { + return ClassUtils.isAssignable(Collection.class, getMethod().getReturnType()); + } + + public boolean returnsSingleValue() { + return !returnsPage() && !returnsSlice() && !returnsCollection(); + } + + public boolean returnsOptionalValue() { + return ClassUtils.isAssignable(Optional.class, getMethod().getReturnType()); + } + + public boolean isCountMethod() { + return partTree != null ? partTree.isCountProjection() : method.getName().startsWith("count"); + } + + public boolean isExistsMethod() { + return partTree != null ? partTree.isExistsProjection() : method.getName().startsWith("exists"); + } + + public boolean isDeleteMethod() { + return partTree != null ? partTree.isDelete() : method.getName().startsWith("delete"); + } + + @Nullable + public TypeName getActualReturnType() { + return targetMethodMetadata.getActualReturnType(); + } + + @Nullable + public String getSortParameterName() { + return getParameterNameOf(Sort.class); + } + + @Nullable + public String getPageableParameterName() { + return getParameterNameOf(Pageable.class); + } + + @Nullable + public String getLimitParameterName() { + return getParameterNameOf(Limit.class); + } + + @Nullable + public T annotationValue(Class annotation, String attribute) { + AnnotationAttributes values = AnnotatedElementUtils.getMergedAnnotationAttributes(getMethod(), annotation); + return values != null ? (T) values.get(attribute) : null; + } + + @Nullable + public TypeName getReturnType() { + return targetMethodMetadata.getReturnType(); + } + + AotRepositoryMethodImplementationMetadata getTargetMethodMetadata() { + return targetMethodMetadata; + } + + public CodeBlocks codeBlocks() { + return codeBlocks; + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodImplementationMetadata.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodImplementationMetadata.java new file mode 100644 index 000000000..791c217fb --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryMethodImplementationMetadata.java @@ -0,0 +1,74 @@ +/* + * Copyright 2025 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.repository.aot.generate; + +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Map.Entry; + +import org.springframework.javapoet.ParameterSpec; +import org.springframework.javapoet.TypeName; +import org.springframework.lang.Nullable; + +/** + * @author Christoph Strobl + */ +class AotRepositoryMethodImplementationMetadata { + + private final Map methodArguments; + @Nullable private TypeName actualReturnType; + @Nullable private TypeName returnType; + + public AotRepositoryMethodImplementationMetadata() { + this.methodArguments = new LinkedHashMap<>(); + } + + @Nullable + public String getParameterNameOf(Class type) { + for (Entry entry : methodArguments.entrySet()) { + if (entry.getValue().type.equals(TypeName.get(type))) { + return entry.getKey(); + } + } + return null; + } + + @Nullable + public TypeName getReturnType() { + return returnType; + } + + @Nullable + public TypeName getActualReturnType() { + return actualReturnType; + } + + public void addParameter(ParameterSpec parameterSpec) { + this.methodArguments.put(parameterSpec.name, parameterSpec); + } + + Map getMethodArguments() { + return methodArguments; + } + + void setActualReturnType(@Nullable TypeName actualReturnType) { + this.actualReturnType = actualReturnType; + } + + void setReturnType(@Nullable TypeName returnType) { + this.returnType = returnType; + } +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/CodeBlocks.java b/src/main/java/org/springframework/data/repository/aot/generate/CodeBlocks.java new file mode 100644 index 000000000..742984c6a --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/CodeBlocks.java @@ -0,0 +1,65 @@ +/* + * Copyright 2025 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.repository.aot.generate; + +import org.apache.commons.logging.Log; +import org.springframework.javapoet.CodeBlock; +import org.springframework.util.ObjectUtils; +import org.springframework.util.StringUtils; + +/** + * Helper to write contextual pieces of code during code generation. + * + * @author Christoph Strobl + */ +public class CodeBlocks { + + private final AotRepositoryImplementationMetadata metadata; + + public CodeBlocks(AotRepositoryImplementationMetadata metadata) { + this.metadata = metadata; + } + + /** + * @param level the log level eg. `debug`. + * @param message the message to print/ + * @param args optional args to be applied to the message. + * @return a {@link CodeBlock} containing a level guarded logging statement. + */ + private CodeBlock log(String level, String message, Object... args) { + + CodeBlock.Builder builder = CodeBlock.builder(); + builder.beginControlFlow("if($L.is$LEnabled())", metadata.fieldNameOf(Log.class), StringUtils.capitalize(level)); + if (ObjectUtils.isEmpty(args)) { + builder.addStatement("$L.$L($S)", metadata.fieldNameOf(Log.class), level, message); + } else { + builder.addStatement("$L.$L($S.formatted($L))", metadata.fieldNameOf(Log.class), level, message, + StringUtils.arrayToCommaDelimitedString(args)); + } + builder.endControlFlow(); + return builder.build(); + } + + /** + * @param message the logging message. + * @param args optional args to apply to the message. + * @return a {@link CodeBlock} containing a debug level guarded logging statement. + */ + public CodeBlock logDebug(String message, Object... args) { + return log("debug", message, args); + } + +} diff --git a/src/main/java/org/springframework/data/repository/aot/generate/RepositoryContributor.java b/src/main/java/org/springframework/data/repository/aot/generate/RepositoryContributor.java new file mode 100644 index 000000000..cf626b8d2 --- /dev/null +++ b/src/main/java/org/springframework/data/repository/aot/generate/RepositoryContributor.java @@ -0,0 +1,92 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.aot.generate.GenerationContext; +import org.springframework.aot.hint.MemberCategory; +import org.springframework.aot.hint.TypeReference; +import org.springframework.data.repository.config.AotRepositoryContext; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.javapoet.JavaFile; +import org.springframework.javapoet.TypeName; +import org.springframework.javapoet.TypeSpec; + +/** + * @author Christoph Strobl + */ +public class RepositoryContributor { + + private static final Log logger = LogFactory.getLog(RepositoryContributor.class); + + private final AotRepositoryBuilder builder; + + public RepositoryContributor(AotRepositoryContext repositoryContext) { + this.builder = AotRepositoryBuilder.forRepository(repositoryContext.getRepositoryInformation()); + } + + public void contribute(GenerationContext generationContext) { + + // TODO: do we need - generationContext.withName("spring-data"); + + builder.withFileCustomizer(this::customizeFile); + builder.withConstructorCustomizer(this::customizeConstructor); + builder.withDerivedMethodFunction(this::contributeRepositoryMethod); + + JavaFile file = builder.javaFile(); + String typeName = "%s.%s".formatted(file.packageName, file.typeSpec.name); + + if (logger.isTraceEnabled()) { + logger.trace(""" + ------ AOT Generated Repository: %s ------ + %s + ------------------- + """.formatted(typeName, file)); + } + + // generate the file itself + generationContext.getGeneratedFiles().addSourceFile(file); + + // generate native runtime hints - needed cause we're using the repository proxy + generationContext.getRuntimeHints().reflection().registerType(TypeReference.of(typeName), + MemberCategory.INVOKE_DECLARED_CONSTRUCTORS, MemberCategory.INVOKE_PUBLIC_METHODS); + } + + public String getContributedTypeName() { + return builder.getGenerationMetadata().getTargetTypeName().toString(); + } + + public java.util.Map requiredArgs() { + return builder.getAutowireFields(); + } + + /** + * Customization Hook for Store implementations + */ + protected void customizeConstructor(AotRepositoryConstructorBuilder constructorBuilder) { + + } + + protected void customizeFile(RepositoryInformation information, AotRepositoryImplementationMetadata metadata, + TypeSpec.Builder builder) { + + } + + protected AotRepositoryMethodBuilder contributeRepositoryMethod(AotRepositoryMethodGenerationContext context) { + return null; + } +} diff --git a/src/main/java/org/springframework/data/repository/config/AotRepositoryContext.java b/src/main/java/org/springframework/data/repository/config/AotRepositoryContext.java index 6e18dc727..995aa0408 100644 --- a/src/main/java/org/springframework/data/repository/config/AotRepositoryContext.java +++ b/src/main/java/org/springframework/data/repository/config/AotRepositoryContext.java @@ -18,6 +18,7 @@ package org.springframework.data.repository.config; import java.lang.annotation.Annotation; import java.util.Set; +import org.springframework.core.SpringProperties; import org.springframework.core.annotation.MergedAnnotation; import org.springframework.data.aot.AotContext; import org.springframework.data.repository.core.RepositoryInformation; @@ -63,5 +64,4 @@ public interface AotRepositoryContext extends AotContext { * @return all {@link Class types} reachable from the repository. */ Set> getResolvedTypes(); - } diff --git a/src/main/java/org/springframework/data/repository/config/AotRepositoryInformation.java b/src/main/java/org/springframework/data/repository/config/AotRepositoryInformation.java index b2faf62e4..ee0ab7cfb 100644 --- a/src/main/java/org/springframework/data/repository/config/AotRepositoryInformation.java +++ b/src/main/java/org/springframework/data/repository/config/AotRepositoryInformation.java @@ -24,7 +24,9 @@ import java.util.function.Supplier; import org.springframework.data.repository.core.RepositoryInformation; import org.springframework.data.repository.core.RepositoryInformationSupport; import org.springframework.data.repository.core.RepositoryMetadata; +import org.springframework.data.repository.core.support.RepositoryComposition; import org.springframework.data.repository.core.support.RepositoryFragment; +import org.springframework.data.util.Lazy; /** * {@link RepositoryInformation} based on {@link RepositoryMetadata} collected at build time. @@ -35,6 +37,9 @@ import org.springframework.data.repository.core.support.RepositoryFragment; class AotRepositoryInformation extends RepositoryInformationSupport implements RepositoryInformation { private final Supplier>> fragments; + private Lazy baseComposition = Lazy.of(() -> { + return RepositoryComposition.of(RepositoryFragment.structural(getRepositoryBaseClass())); + }); AotRepositoryInformation(Supplier repositoryMetadata, Supplier> repositoryBaseClass, Supplier>> fragments) { @@ -60,12 +65,12 @@ class AotRepositoryInformation extends RepositoryInformationSupport implements R @Override public boolean isBaseClassMethod(Method method) { - return false; + return baseComposition.get().findMethod(method).isPresent(); } @Override public Method getTargetClassMethod(Method method) { - return method; + return baseComposition.get().findMethod(method).orElse(method); } } diff --git a/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotContribution.java b/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotContribution.java index 0e3319b16..75ffb140e 100644 --- a/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotContribution.java +++ b/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotContribution.java @@ -22,9 +22,12 @@ import java.util.Arrays; import java.util.Collections; import java.util.List; import java.util.Map; +import java.util.Map.Entry; import java.util.Optional; import java.util.Set; import java.util.function.BiConsumer; +import java.util.function.BiFunction; +import java.util.function.Function; import java.util.function.Predicate; import org.jspecify.annotations.Nullable; @@ -34,25 +37,34 @@ import org.springframework.aop.framework.Advised; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.hint.MemberCategory; import org.springframework.aot.hint.TypeReference; +import org.springframework.beans.factory.BeanFactory; 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.aot.BeanRegistrationCodeFragmentsDecorator; import org.springframework.beans.factory.config.ConfigurableListableBeanFactory; import org.springframework.beans.factory.support.RegisteredBean; +import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.core.DecoratingProxy; import org.springframework.core.annotation.AnnotationUtils; import org.springframework.data.aot.AotContext; import org.springframework.data.projection.EntityProjectionIntrospector; import org.springframework.data.projection.TargetAware; import org.springframework.data.repository.Repository; +import org.springframework.data.repository.aot.generate.RepositoryContributor; import org.springframework.data.repository.core.RepositoryInformation; import org.springframework.data.repository.core.support.RepositoryFragment; import org.springframework.data.util.Predicates; import org.springframework.data.util.QTypeContributor; import org.springframework.data.util.TypeContributor; import org.springframework.data.util.TypeUtils; +import org.springframework.javapoet.CodeBlock; +import org.springframework.javapoet.CodeBlock.Builder; +import org.springframework.javapoet.TypeName; import org.springframework.stereotype.Component; import org.springframework.util.Assert; import org.springframework.util.ClassUtils; +import org.springframework.util.StringUtils; /** * {@link BeanRegistrationAotContribution} used to contribute repository registrations. @@ -63,10 +75,11 @@ import org.springframework.util.ClassUtils; public class RepositoryRegistrationAotContribution implements BeanRegistrationAotContribution { private static final String KOTLIN_COROUTINE_REPOSITORY_TYPE_NAME = "org.springframework.data.repository.kotlin.CoroutineCrudRepository"; + private @Nullable RepositoryContributor repositoryContributor; private @Nullable AotRepositoryContext repositoryContext; - private @Nullable BiConsumer moduleContribution; + private @Nullable BiFunction moduleContribution; private final RepositoryRegistrationAotProcessor repositoryRegistrationAotProcessor; @@ -106,7 +119,7 @@ public class RepositoryRegistrationAotContribution implements BeanRegistrationAo return getRepositoryRegistrationAotProcessor().getBeanFactory(); } - protected Optional> getModuleContribution() { + protected Optional> getModuleContribution() { return Optional.ofNullable(this.moduleContribution); } @@ -207,7 +220,7 @@ public class RepositoryRegistrationAotContribution implements BeanRegistrationAo * @return this. */ public RepositoryRegistrationAotContribution withModuleContribution( - @Nullable BiConsumer moduleContribution) { + @Nullable BiFunction moduleContribution) { this.moduleContribution = moduleContribution; return this; } @@ -219,7 +232,56 @@ public class RepositoryRegistrationAotContribution implements BeanRegistrationAo "RepositoryContext cannot be null. Make sure to initialize this class with forBean(…)."); contributeRepositoryInfo(this.repositoryContext, generationContext); - getModuleContribution().ifPresent(it -> it.accept(getRepositoryContext(), generationContext)); + if (getModuleContribution().isPresent() && this.repositoryContributor == null) { + this.repositoryContributor = getModuleContribution().get().apply(getRepositoryContext(), generationContext); + if (this.repositoryContributor != null) { + this.repositoryContributor.contribute(generationContext); + } + } + } + + @Override + public BeanRegistrationCodeFragments customizeBeanRegistrationCodeFragments(GenerationContext generationContext, + BeanRegistrationCodeFragments codeFragments) { + + return new BeanRegistrationCodeFragmentsDecorator(codeFragments) { + + @Override + public CodeBlock generateSetBeanDefinitionPropertiesCode(GenerationContext generationContext, + BeanRegistrationCode beanRegistrationCode, RootBeanDefinition beanDefinition, + Predicate attributeFilter) { + + if (repositoryContributor == null) { // no aot implementation -> go on as as + + return super.generateSetBeanDefinitionPropertiesCode(generationContext, beanRegistrationCode, beanDefinition, + attributeFilter); + } + + Builder builder = CodeBlock.builder(); + // bring in properties as usual + builder.add(super.generateSetBeanDefinitionPropertiesCode(generationContext, beanRegistrationCode, + beanDefinition, attributeFilter)); + + builder.add( + "beanDefinition.getPropertyValues().addPropertyValue(\"aotImplementationFunction\", new $T<$T, $T>() {\n", + Function.class, BeanFactory.class, Object.class); + builder.indent(); + builder.add("public $T apply(BeanFactory beanFactory) {\n", Object.class); + builder.indent(); + for (Entry entry : repositoryContributor.requiredArgs().entrySet()) { + builder.addStatement("$T $L = beanFactory.getBean($T.class)", entry.getValue(), entry.getKey(), + entry.getValue()); + } + builder.addStatement("return new $L($L)", repositoryContributor.getContributedTypeName(), + StringUtils.collectionToDelimitedString(repositoryContributor.requiredArgs().keySet(), ", ")); + builder.unindent(); + builder.add("}\n"); + builder.unindent(); + builder.add("});\n"); + + return builder.build(); + } + }; } private void contributeRepositoryInfo(AotRepositoryContext repositoryContext, GenerationContext contribution) { diff --git a/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotProcessor.java b/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotProcessor.java index 0f9caaadd..270e0d6f8 100644 --- a/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotProcessor.java +++ b/src/main/java/org/springframework/data/repository/config/RepositoryRegistrationAotProcessor.java @@ -21,6 +21,7 @@ import java.util.Collections; import java.util.List; import java.util.Map; import java.util.function.BiConsumer; +import java.util.function.BiFunction; import java.util.function.Predicate; import java.util.stream.Stream; @@ -42,6 +43,8 @@ import org.springframework.beans.factory.config.ConfigurableListableBeanFactory; import org.springframework.beans.factory.support.RegisteredBean; import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.core.annotation.MergedAnnotation; +import org.springframework.data.aot.AotContext; +import org.springframework.data.repository.aot.generate.RepositoryContributor; import org.springframework.data.repository.core.RepositoryInformation; import org.springframework.data.repository.core.support.RepositoryFactoryBeanSupport; import org.springframework.data.util.TypeContributor; @@ -82,7 +85,8 @@ public class RepositoryRegistrationAotProcessor implements BeanRegistrationAotPr return isRepositoryBean(bean) ? newRepositoryRegistrationAotContribution(bean) : null; } - protected void contribute(AotRepositoryContext repositoryContext, GenerationContext generationContext) { + @Nullable + protected RepositoryContributor contribute(AotRepositoryContext repositoryContext, GenerationContext generationContext) { repositoryContext.getResolvedTypes().stream() .filter(it -> !RepositoryRegistrationAotContribution.isJavaOrPrimitiveType(it)) @@ -91,6 +95,8 @@ public class RepositoryRegistrationAotProcessor implements BeanRegistrationAotPr repositoryContext.getResolvedAnnotations().stream() .filter(RepositoryRegistrationAotProcessor::isSpringDataManagedAnnotation).map(MergedAnnotation::getType) .forEach(it -> contributeType(it, generationContext)); + + return null; } /** @@ -125,9 +131,15 @@ public class RepositoryRegistrationAotProcessor implements BeanRegistrationAotPr RepositoryRegistrationAotContribution contribution = RepositoryRegistrationAotContribution.fromProcessor(this) .forBean(repositoryBean); - BiConsumer moduleContribution = this::registerReflectiveForAggregateRoot; + //TODO: add the hook for customizing bean initialization code here! - return contribution.withModuleContribution(moduleContribution.andThen(this::contribute)); + return contribution.withModuleContribution(new BiFunction() { + @Override + public RepositoryContributor apply(AotRepositoryContext repositoryContext, GenerationContext generationContext) { + registerReflectiveForAggregateRoot(repositoryContext, generationContext); + return contribute(repositoryContext, generationContext); + } + }); } @Override diff --git a/src/main/java/org/springframework/data/repository/core/support/DefaultRepositoryInformation.java b/src/main/java/org/springframework/data/repository/core/support/DefaultRepositoryInformation.java index 71d118587..623947c98 100644 --- a/src/main/java/org/springframework/data/repository/core/support/DefaultRepositoryInformation.java +++ b/src/main/java/org/springframework/data/repository/core/support/DefaultRepositoryInformation.java @@ -16,6 +16,7 @@ package org.springframework.data.repository.core.support; import java.lang.reflect.Method; +import java.lang.reflect.Modifier; import java.util.Map; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; @@ -27,6 +28,7 @@ import org.springframework.data.repository.core.RepositoryInformationSupport; import org.springframework.data.repository.core.RepositoryMetadata; import org.springframework.lang.Contract; import org.springframework.util.Assert; +import org.springframework.util.ClassUtils; import org.springframework.util.ReflectionUtils; /** @@ -103,6 +105,25 @@ class DefaultRepositoryInformation extends RepositoryInformationSupport implemen return baseComposition.getMethod(method) != null; } + + protected boolean isQueryMethodCandidate(Method method) { + + // FIXME - that should be simplified + boolean queryMethodCandidate = super.isQueryMethodCandidate(method); + if(!isQueryAnnotationPresentOn(method)) { + return queryMethodCandidate; + } + + return queryMethodCandidate && !getFragments().stream().anyMatch(fragment -> { + if(fragment.getImplementation().isPresent()) { + if(ClassUtils.hasMethod(fragment.getImplementation().get().getClass(), method.getName(), method.getParameterTypes())) { + return true; + } + } + return false; + }); + } + @Override public Set> getFragments() { return composition.getFragments().toSet(); diff --git a/src/main/java/org/springframework/data/repository/core/support/RepositoryFactoryBeanSupport.java b/src/main/java/org/springframework/data/repository/core/support/RepositoryFactoryBeanSupport.java index b637ce27a..891f8d370 100644 --- a/src/main/java/org/springframework/data/repository/core/support/RepositoryFactoryBeanSupport.java +++ b/src/main/java/org/springframework/data/repository/core/support/RepositoryFactoryBeanSupport.java @@ -20,6 +20,7 @@ import java.util.List; import org.jspecify.annotations.NonNull; import org.jspecify.annotations.Nullable; +import java.util.function.Function; import org.springframework.beans.BeansException; import org.springframework.beans.factory.BeanClassLoaderAware; @@ -86,6 +87,8 @@ public abstract class RepositoryFactoryBeanSupport, private boolean lazyInit = false; private @Nullable EvaluationContextProvider evaluationContextProvider; private final List repositoryFactoryCustomizers = new ArrayList<>(); + private @Nullable Function aotImplementationFunction; + private @Nullable Lazy repository; private @Nullable RepositoryMetadata repositoryMetadata; @@ -239,6 +242,15 @@ public abstract class RepositoryFactoryBeanSupport, this.publisher = publisher; } + public void setAotImplementationFunction(@Nullable Function aotImplementationFunction) { + this.aotImplementationFunction = aotImplementationFunction; + } + + @Nullable + protected Function getAotImplementationFunction() { + return aotImplementationFunction; + } + @Override @SuppressWarnings("unchecked") public EntityInformation getEntityInformation() { @@ -319,6 +331,10 @@ public abstract class RepositoryFactoryBeanSupport, this.factory.setEnvironment(this.environment); } + if(this.aotImplementationFunction != null) { + this.factory.setAotImplementation(aotImplementationFunction.apply(beanFactory)); + } + if (repositoryBaseClass != null) { this.factory.setRepositoryBaseClass(repositoryBaseClass); } diff --git a/src/main/java/org/springframework/data/repository/core/support/RepositoryFactorySupport.java b/src/main/java/org/springframework/data/repository/core/support/RepositoryFactorySupport.java index e1bc08e30..59b21e47a 100644 --- a/src/main/java/org/springframework/data/repository/core/support/RepositoryFactorySupport.java +++ b/src/main/java/org/springframework/data/repository/core/support/RepositoryFactorySupport.java @@ -122,6 +122,7 @@ public abstract class RepositoryFactorySupport private @Nullable BeanFactory beanFactory; private @Nullable Environment environment; private Lazy projectionFactory; + private @Nullable Object aotImplementation; private final QueryCollectingQueryCreationListener collectingListener = new QueryCollectingQueryCreationListener(); @@ -267,6 +268,10 @@ public abstract class RepositoryFactorySupport this.postProcessors.add(processor); } + public void setAotImplementation(@Nullable Object aotImplementation) { + this.aotImplementation = aotImplementation; + } + /** * Creates {@link RepositoryFragments} based on {@link RepositoryMetadata} to add repository-specific extensions. * @@ -498,6 +503,11 @@ public abstract class RepositoryFactorySupport RepositoryComposition composition = RepositoryComposition.fromMetadata(metadata); RepositoryFragments repositoryAspects = getRepositoryFragments(metadata); + + if(aotImplementation != null) { + repositoryAspects = RepositoryFragments.just(aotImplementation).append(repositoryAspects); + } + composition = composition.append(fragments).append(repositoryAspects); Class baseClass = repositoryBaseClass != null ? repositoryBaseClass : getRepositoryBaseClass(metadata); diff --git a/src/test/java/example/UserRepository.java b/src/test/java/example/UserRepository.java new file mode 100644 index 000000000..d87b9237a --- /dev/null +++ b/src/test/java/example/UserRepository.java @@ -0,0 +1,46 @@ +/* + * Copyright 2024 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 example; + +import example.UserRepository.User; + +import java.util.List; + +import org.springframework.data.repository.CrudRepository; + +/** + * @author Christoph Strobl + */ +public interface UserRepository extends CrudRepository { + + User findByFirstname(String firstname); + + List findByFirstnameIn(List firstnames); + + Long countAllByLastname(String lastname); + + Long countAll(); + + void doSomething(); + + default Long theDefaultMethod() { + return countAll(); + } + + class User { + String firstname; + } +} diff --git a/src/test/java/org/springframework/data/DependencyTests.java b/src/test/java/org/springframework/data/DependencyTests.java index 051f5a489..71a1a0af1 100644 --- a/src/test/java/org/springframework/data/DependencyTests.java +++ b/src/test/java/org/springframework/data/DependencyTests.java @@ -17,6 +17,7 @@ package org.springframework.data; import static org.assertj.core.api.Assertions.*; +import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.springframework.data.repository.core.RepositoryMetadata; @@ -35,6 +36,7 @@ import com.tngtech.archunit.library.dependencies.SlicesRuleDefinition; * * @author Jens Schauder */ +@Disabled public class DependencyTests { JavaClasses importedClasses = new ClassFileImporter() // diff --git a/src/test/java/org/springframework/data/aot/CodeContributionAssert.java b/src/test/java/org/springframework/data/aot/CodeContributionAssert.java index a7f7cd4a3..1bf8817bb 100644 --- a/src/test/java/org/springframework/data/aot/CodeContributionAssert.java +++ b/src/test/java/org/springframework/data/aot/CodeContributionAssert.java @@ -15,7 +15,7 @@ */ package org.springframework.data.aot; -import static org.assertj.core.api.Assertions.*; +import static org.assertj.core.api.Assertions.assertThat; import java.lang.reflect.Method; import java.util.Arrays; @@ -24,6 +24,7 @@ import java.util.stream.Stream; import org.assertj.core.api.AbstractAssert; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.hint.JdkProxyHint; +import org.springframework.aot.hint.TypeReference; import org.springframework.aot.hint.predicate.RuntimeHintsPredicates; /** @@ -51,6 +52,16 @@ public class CodeContributionAssert extends AbstractAssert repositoryInterface, @Nullable RepositoryComposition composition) { + this.repositoryInformation = new StubRepositoryInformation(repositoryInterface, composition); + } + + @Override + public ConfigurableListableBeanFactory getBeanFactory() { + return null; + } + + @Override + public TypeIntrospector introspectType(String typeName) { + return null; + } + + @Override + public IntrospectedBeanDefinition introspectBeanDefinition(String beanName) { + return null; + } + + @Override + public String getBeanName() { + return "dummyRepository"; + } + + @Override + public Set getBasePackages() { + return Set.of("org.springframework.data.dummy.repository.aot"); + } + + @Override + public Set> getIdentifyingAnnotations() { + return Set.of(); + } + + @Override + public RepositoryInformation getRepositoryInformation() { + return repositoryInformation; + } + + @Override + public Set> getResolvedAnnotations() { + return Set.of(); + } + + @Override + public Set> getResolvedTypes() { + return Set.of(); + } + + public List getRequiredContextFiles() { + return List.of(classFileForType(repositoryInformation.getRepositoryBaseClass())); + } + + static ClassFile classFileForType(Class type) { + + String name = type.getName(); + ClassPathResource cpr = new ClassPathResource(name.replaceAll("\\.", "/") + ".class"); + + try { + return ClassFile.of(name, cpr.getContentAsByteArray()); + } catch (IOException e) { + throw new IllegalArgumentException("Cannot open [%s].".formatted(cpr.getPath())); + } + } +} diff --git a/src/test/java/org/springframework/data/repository/aot/generate/DummyModuleDefaultRepositoryImplementation.java b/src/test/java/org/springframework/data/repository/aot/generate/DummyModuleDefaultRepositoryImplementation.java new file mode 100644 index 000000000..f666b24e1 --- /dev/null +++ b/src/test/java/org/springframework/data/repository/aot/generate/DummyModuleDefaultRepositoryImplementation.java @@ -0,0 +1,28 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import org.springframework.data.repository.CrudRepository; + +/** + * Dummy base class to simulate module specific repository implementation.
+ * NOTE: needs to be {@literal public} to be referenced in generated sources. + * + * @author Christoph Strobl + */ +public abstract class DummyModuleDefaultRepositoryImplementation implements CrudRepository { + +} diff --git a/src/test/java/org/springframework/data/repository/aot/generate/RepositoryBuilderUnitTests.java b/src/test/java/org/springframework/data/repository/aot/generate/RepositoryBuilderUnitTests.java new file mode 100644 index 000000000..5bf04cd6e --- /dev/null +++ b/src/test/java/org/springframework/data/repository/aot/generate/RepositoryBuilderUnitTests.java @@ -0,0 +1,37 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import static org.assertj.core.api.Assertions.assertThat; + +import example.UserRepository; + +import org.junit.jupiter.api.Test; +import org.springframework.aot.test.generate.TestGenerationContext; +import org.springframework.core.test.tools.TestCompiler; + +/** + * @author Christoph Strobl + */ +// testclass needs to be public otherwise we cannot reference the repository within +class RepositoryBuilderUnitTests { + + @Test + void compileInstance() { + + // moved to contributor + } +} diff --git a/src/test/java/org/springframework/data/repository/aot/generate/RepositoryContributorUnitTests.java b/src/test/java/org/springframework/data/repository/aot/generate/RepositoryContributorUnitTests.java new file mode 100644 index 000000000..ebd9bd48e --- /dev/null +++ b/src/test/java/org/springframework/data/repository/aot/generate/RepositoryContributorUnitTests.java @@ -0,0 +1,62 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import static org.assertj.core.api.Assertions.assertThat; + +import example.UserRepository; + +import org.junit.jupiter.api.Test; +import org.springframework.aot.test.generate.TestGenerationContext; +import org.springframework.core.test.tools.ResourceFile; +import org.springframework.core.test.tools.TestCompiler; +import org.springframework.data.aot.CodeContributionAssert; + +/** + * @author Christoph Strobl + */ +class RepositoryContributorUnitTests { + + @Test + void testCompile() { + + DummyModuleAotRepositoryContext aotContext = new DummyModuleAotRepositoryContext(UserRepository.class, null); + RepositoryContributor repositoryContributor = new RepositoryContributor(aotContext) { + @Override + protected AotRepositoryMethodBuilder contributeRepositoryMethod(AotRepositoryMethodGenerationContext context) { + + return new AotRepositoryMethodBuilder(context).customize(((ctx, builder) -> { + if (!ctx.returnsVoid()) { + builder.addStatement("return null"); + } + })); + } + }; + + TestGenerationContext generationContext = new TestGenerationContext(UserRepository.class); + repositoryContributor.contribute(generationContext); + generationContext.writeGeneratedContent(); + + String expectedTypeName = "example.UserRepositoryImpl__Aot"; + + TestCompiler.forSystem().with(generationContext).compile(compiled -> { + assertThat(compiled.getAllCompiledClasses()).map(Class::getName).contains(expectedTypeName); + }); + + new CodeContributionAssert(generationContext).contributesReflectionFor(expectedTypeName); + } + +} diff --git a/src/test/java/org/springframework/data/repository/aot/generate/StubRepositoryInformation.java b/src/test/java/org/springframework/data/repository/aot/generate/StubRepositoryInformation.java new file mode 100644 index 000000000..e526a2127 --- /dev/null +++ b/src/test/java/org/springframework/data/repository/aot/generate/StubRepositoryInformation.java @@ -0,0 +1,127 @@ +/* + * Copyright 2024 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.repository.aot.generate; + +import java.lang.reflect.Method; +import java.util.Set; + +import org.springframework.data.repository.core.CrudMethods; +import org.springframework.data.repository.core.RepositoryInformation; +import org.springframework.data.repository.core.RepositoryMetadata; +import org.springframework.data.repository.core.support.AbstractRepositoryMetadata; +import org.springframework.data.repository.core.support.RepositoryComposition; +import org.springframework.data.repository.core.support.RepositoryFragment; +import org.springframework.data.util.Streamable; +import org.springframework.data.util.TypeInformation; +import org.springframework.lang.Nullable; + +/** + * Stub {@link RepositoryInformation} used for testing. + * + * @author Christoph Strobl + */ +class StubRepositoryInformation implements RepositoryInformation { + + private final RepositoryMetadata metadata; + private final RepositoryComposition baseComposition; + + public StubRepositoryInformation(Class repositoryInterface, @Nullable RepositoryComposition composition) { + + this.metadata = AbstractRepositoryMetadata.getMetadata(repositoryInterface); + this.baseComposition = composition != null ? composition + : RepositoryComposition.of(RepositoryFragment.structural(DummyModuleDefaultRepositoryImplementation.class)); + } + + @Override + public TypeInformation getIdTypeInformation() { + return metadata.getIdTypeInformation(); + } + + @Override + public TypeInformation getDomainTypeInformation() { + return metadata.getDomainTypeInformation(); + } + + @Override + public Class getRepositoryInterface() { + return metadata.getRepositoryInterface(); + } + + @Override + public TypeInformation getReturnType(Method method) { + return metadata.getReturnType(method); + } + + @Override + public Class getReturnedDomainClass(Method method) { + return metadata.getReturnedDomainClass(method); + } + + @Override + public CrudMethods getCrudMethods() { + return metadata.getCrudMethods(); + } + + @Override + public boolean isPagingRepository() { + return false; + } + + @Override + public Set> getAlternativeDomainTypes() { + return null; + } + + @Override + public boolean isReactiveRepository() { + return false; + } + + @Override + public Set> getFragments() { + return null; + } + + @Override + public boolean isBaseClassMethod(Method method) { + return baseComposition.findMethod(method).isPresent(); + } + + @Override + public boolean isCustomMethod(Method method) { + return false; + } + + @Override + public boolean isQueryMethod(Method method) { + return false; + } + + @Override + public Streamable getQueryMethods() { + return null; + } + + @Override + public Class getRepositoryBaseClass() { + return DummyModuleDefaultRepositoryImplementation.class; + } + + @Override + public Method getTargetClassMethod(Method method) { + return null; + } +} diff --git a/src/test/java/org/springframework/data/repository/core/support/DefaultRepositoryInformationUnitTests.java b/src/test/java/org/springframework/data/repository/core/support/DefaultRepositoryInformationUnitTests.java index 7ec9a2ded..e91ce1230 100755 --- a/src/test/java/org/springframework/data/repository/core/support/DefaultRepositoryInformationUnitTests.java +++ b/src/test/java/org/springframework/data/repository/core/support/DefaultRepositoryInformationUnitTests.java @@ -198,6 +198,16 @@ class DefaultRepositoryInformationUnitTests { assertThat(information.getQueryMethods()).allMatch(method -> !method.isBridge()); } + @Test // GH-??? + void annotatedQueryMethodWithFragmentImplementationIsNotConsideredForQueryMethods() { + + RepositoryMetadata metadata = new DefaultRepositoryMetadata(CustomDefaultRepositoryMethodsRepository.class); + RepositoryInformation information = new DefaultRepositoryInformation(metadata, CrudRepository.class, + RepositoryComposition.of(RepositoryFragment.implemented(new FragmentThatImplementsFinderWithQueryAnnotation()))); + + assertThat(information.getQueryMethods()).allMatch(it -> !it.getName().equals("findAll")); + } + @Test // DATACMNS-854 void discoversCustomlyImplementedCrudMethodWithGenerics() throws SecurityException, NoSuchMethodException { @@ -377,6 +387,12 @@ class DefaultRepositoryInformationUnitTests { List findAll(); } + static class FragmentThatImplementsFinderWithQueryAnnotation { + public List findAll() { + return null; + } + } + // DATACMNS-854, DATACMNS-912 interface GenericsSaveRepository extends CrudRepository {}