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 index 2fbb700fe..db35dbe35 100644 --- a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryBuilder.java @@ -19,6 +19,7 @@ import java.lang.reflect.Method; import java.util.ArrayList; import java.util.Arrays; import java.util.Comparator; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.function.BiFunction; @@ -29,10 +30,10 @@ import javax.lang.model.element.Modifier; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.jspecify.annotations.Nullable; - import org.springframework.aot.generate.ClassNameGenerator; import org.springframework.aot.generate.Generated; import org.springframework.data.projection.ProjectionFactory; +import org.springframework.data.repository.aot.generate.AotRepositoryFragmentMetadata.ConstructorArgument; import org.springframework.data.repository.aot.generate.json.JSONException; import org.springframework.data.repository.aot.generate.json.JSONObject; import org.springframework.data.repository.core.RepositoryInformation; @@ -215,7 +216,11 @@ class AotRepositoryBuilder { } public Map getAutowireFields() { - return generationMetadata.getConstructorArguments(); + Map autowireFields = new LinkedHashMap<>(generationMetadata.getConstructorArguments().size()); + for (Map.Entry entry : generationMetadata.getConstructorArguments().entrySet()) { + autowireFields.put(entry.getKey(), entry.getValue().getTypeName()); + } + return autowireFields; } public RepositoryInformation getRepositoryInformation() { @@ -238,8 +243,7 @@ class AotRepositoryBuilder { * @param metadata * @param builder */ - void customize(RepositoryInformation information, AotRepositoryFragmentMetadata metadata, - TypeSpec.Builder builder); + void customize(RepositoryInformation information, AotRepositoryFragmentMetadata 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 index 0a62a6a5c..0f44e89e8 100644 --- a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryConstructorBuilder.java @@ -21,6 +21,7 @@ import java.util.Map.Entry; import javax.lang.model.element.Modifier; import org.springframework.core.ResolvableType; +import org.springframework.data.repository.aot.generate.AotRepositoryFragmentMetadata.ConstructorArgument; import org.springframework.data.repository.core.RepositoryInformation; import org.springframework.javapoet.MethodSpec; import org.springframework.javapoet.ParameterizedTypeName; @@ -64,15 +65,27 @@ public class AotRepositoryConstructorBuilder { } /** - * Add constructor parameter. + * Add constructor parameter and create a field for it. * * @param parameterName * @param type */ public void addParameter(String parameterName, TypeName type) { + addParameter(parameterName, type, true); + } - this.metadata.addConstructorArgument(parameterName, type); - this.metadata.addField(parameterName, type, Modifier.PRIVATE, Modifier.FINAL); + /** + * Add constructor parameter. + * + * @param parameterName + * @param type + */ + public void addParameter(String parameterName, TypeName type, boolean createField) { + + this.metadata.addConstructorArgument(parameterName, type, createField ? parameterName : null); + if(createField) { + this.metadata.addField(parameterName, type, Modifier.PRIVATE, Modifier.FINAL); + } } /** @@ -89,15 +102,17 @@ public class AotRepositoryConstructorBuilder { MethodSpec.Builder builder = MethodSpec.constructorBuilder().addModifiers(Modifier.PUBLIC); - for (Entry parameter : this.metadata.getConstructorArguments().entrySet()) { - builder.addParameter(parameter.getValue(), parameter.getKey()); + for (Entry parameter : this.metadata.getConstructorArguments().entrySet()) { + builder.addParameter(parameter.getValue().getTypeName(), parameter.getKey()); } customizer.customize(repositoryInformation, builder); - for (Entry parameter : this.metadata.getConstructorArguments().entrySet()) { - builder.addStatement("this.$N = $N", parameter.getKey(), + for (Entry parameter : this.metadata.getConstructorArguments().entrySet()) { + if(parameter.getValue().isForLocalField()) { + builder.addStatement("this.$N = $N", parameter.getKey(), parameter.getKey()); + } } return builder.build(); diff --git a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryFragmentMetadata.java b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryFragmentMetadata.java index cf66c0a05..d0286081c 100644 --- a/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryFragmentMetadata.java +++ b/src/main/java/org/springframework/data/repository/aot/generate/AotRepositoryFragmentMetadata.java @@ -35,7 +35,7 @@ public class AotRepositoryFragmentMetadata { private final ClassName className; private final Map fields = new HashMap<>(3); - private final Map constructorArguments = new LinkedHashMap<>(3); + private final Map constructorArguments = new LinkedHashMap<>(3); public AotRepositoryFragmentMetadata(ClassName className) { this.className = className; @@ -82,11 +82,43 @@ public class AotRepositoryFragmentMetadata { return fields; } - public Map getConstructorArguments() { + public Map getConstructorArguments() { return constructorArguments; } public void addConstructorArgument(String parameterName, TypeName type) { - this.constructorArguments.put(parameterName, type); + addConstructorArgument(parameterName, type, parameterName); + } + + public void addConstructorArgument(String parameterName, TypeName type, @Nullable String fieldName) { + this.constructorArguments.put(parameterName, new ConstructorArgument(parameterName, type, fieldName)); + } + + static class ConstructorArgument { + String parameterName; + @Nullable String fieldName; + TypeName typeName; + + public ConstructorArgument(String parameterName,TypeName typeName, String fieldName) { + this.parameterName = parameterName; + this.fieldName = fieldName; + this.typeName = typeName; + } + + boolean isForLocalField() { + return fieldName != null; + } + + public String getParameterName() { + return parameterName; + } + + public String getFieldName() { + return fieldName; + } + + public TypeName getTypeName() { + return typeName; + } } }