From 802a8db1d1472e208dff75c21dfd8dff884918a8 Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Wed, 26 Mar 2025 08:56:51 +0100 Subject: [PATCH] Add query hint support. See #3830 --- .../aot/generated/JpaCodeBlocks.java | 399 ++++++++++-------- .../generated/JpaRepositoryContributor.java | 4 +- ...RepositoryContributorIntegrationTests.java | 9 +- .../aot/generated/UserRepository.java | 6 + 4 files changed, 241 insertions(+), 177 deletions(-) diff --git a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaCodeBlocks.java b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaCodeBlocks.java index 7bfdb0717..3f249fdf4 100644 --- a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaCodeBlocks.java +++ b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaCodeBlocks.java @@ -17,13 +17,16 @@ package org.springframework.data.jpa.repository.aot.generated; import jakarta.persistence.EntityManager; import jakarta.persistence.Query; +import jakarta.persistence.QueryHint; import java.util.List; import java.util.Optional; import java.util.function.LongSupplier; import java.util.regex.Pattern; +import org.springframework.core.annotation.MergedAnnotation; import org.springframework.data.domain.SliceImpl; +import org.springframework.data.jpa.repository.QueryHints; import org.springframework.data.jpa.repository.query.DeclaredQuery; import org.springframework.data.jpa.repository.query.ParameterBinding; import org.springframework.data.repository.aot.generate.AotRepositoryMethodGenerationContext; @@ -37,36 +40,248 @@ import org.springframework.util.StringUtils; /** * @author Christoph Strobl - * @since 2025/01 + * @author Mark Paluch + * @since 4.0 */ -public class JpaCodeBlocks { +class JpaCodeBlocks { private static final Pattern PARAMETER_BINDING_PATTERN = Pattern.compile("\\?(\\d+)"); - static QueryBlockBuilder queryBlockBuilder(AotRepositoryMethodGenerationContext context) { + public static QueryBlockBuilder queryBuilder(AotRepositoryMethodGenerationContext context) { return new QueryBlockBuilder(context); } - static QueryExecutionBlockBuilder queryExecutionBlockBuilder(AotRepositoryMethodGenerationContext context) { + static QueryExecutionBlockBuilder executionBuilder(AotRepositoryMethodGenerationContext context) { return new QueryExecutionBlockBuilder(context); } - static class QueryExecutionBlockBuilder { + /** + * Builder for the actual query code block. + */ + static class QueryBlockBuilder { - AotRepositoryMethodGenerationContext context; + private final AotRepositoryMethodGenerationContext context; private String queryVariableName = "query"; + private AotQueries queries; + private MergedAnnotation queryHints = MergedAnnotation.missing(); - public QueryExecutionBlockBuilder(AotRepositoryMethodGenerationContext context) { + private QueryBlockBuilder(AotRepositoryMethodGenerationContext context) { this.context = context; } - QueryExecutionBlockBuilder referencing(String queryVariableName) { + public QueryBlockBuilder usingQueryVariableName(String queryVariableName) { this.queryVariableName = queryVariableName; return this; } - CodeBlock build() { + public QueryBlockBuilder filter(AotQueries query) { + this.queries = query; + return this; + } + + public QueryBlockBuilder queryHints(MergedAnnotation queryHints) { + + this.queryHints = queryHints; + return this; + } + + /** + * Build the query block. + * + * @return + */ + public CodeBlock build() { + + boolean isProjecting = context.getActualReturnType() != null + && !ObjectUtils.nullSafeEquals(TypeName.get(context.getRepositoryInformation().getDomainType()), + context.getActualReturnType()); + Object actualReturnType = isProjecting ? context.getActualReturnType() + : context.getRepositoryInformation().getDomainType(); + + CodeBlock.Builder builder = CodeBlock.builder(); + builder.add("\n"); + String queryStringNameVariableName = "%sString".formatted(queryVariableName); + + StringAotQuery query = (StringAotQuery) queries.result(); + builder.addStatement("$T $L = $S", String.class, queryStringNameVariableName, query.getQueryString()); + + String countQueryStringNameVariableName = null; + String countQuyerVariableName = null; + + if (context.returnsPage()) { + + countQueryStringNameVariableName = "count%sString".formatted(StringUtils.capitalize(queryVariableName)); + countQuyerVariableName = "count%s".formatted(StringUtils.capitalize(queryVariableName)); + + StringAotQuery countQuery = (StringAotQuery) queries.count(); + builder.addStatement("$T $L = $S", String.class, countQueryStringNameVariableName, countQuery.getQueryString()); + } + + // sorting + // TODO: refactor into sort builder + + String sortParameterName = context.getSortParameterName(); + if (sortParameterName == null && context.getPageableParameterName() != null) { + sortParameterName = "%s.getSort()".formatted(context.getPageableParameterName()); + } + + if (StringUtils.hasText(sortParameterName)) { + builder.add(applySorting(sortParameterName, queryStringNameVariableName, actualReturnType)); + } + + builder.add(createQuery(queryVariableName, queryStringNameVariableName, queries.result(), queryHints)); + + builder.add(applyLimits()); + + if (StringUtils.hasText(countQueryStringNameVariableName)) { + + builder.beginControlFlow("$T $L = () ->", LongSupplier.class, "countAll"); + + boolean queryHints = this.queryHints.isPresent() && this.queryHints.getBoolean("forCounting"); + + builder.add(createQuery(countQuyerVariableName, countQueryStringNameVariableName, queries.count(), + queryHints ? this.queryHints : MergedAnnotation.missing())); + builder.addStatement("return ($T) $L.getSingleResult()", Long.class, countQuyerVariableName); + + // end control flow does not work well with lambdas + builder.unindent(); + builder.add("};\n"); + } + + return builder.build(); + } + + private CodeBlock applySorting(String sort, String queryString, Object actualReturnType) { + + Builder builder = CodeBlock.builder(); + builder.beginControlFlow("if ($L.isSorted())", sort); + + if (queries.isNative()) { + builder.addStatement("$T declaredQuery = $T.nativeQuery($L)", DeclaredQuery.class, DeclaredQuery.class, + queryString); + } else { + builder.addStatement("$T declaredQuery = $T.jpqlQuery($L)", DeclaredQuery.class, DeclaredQuery.class, + queryString); + } + + builder.addStatement("$L = rewriteQuery(declaredQuery, $L, $T.class)", queryString, sort, actualReturnType); + + builder.endControlFlow(); + + return builder.build(); + } + + private CodeBlock applyLimits() { + + Builder builder = CodeBlock.builder(); + + if (context.isExistsMethod()) { + builder.addStatement("$L.setMaxResults(1)", queryVariableName); + + return builder.build(); + } + + String limit = context.getLimitParameterName(); + + if (StringUtils.hasText(limit)) { + builder.beginControlFlow("if ($L.isLimited())", limit); + builder.addStatement("$L.setMaxResults($L.max())", queryVariableName, limit); + builder.endControlFlow(); + } else if (queries.result().isLimited()) { + builder.addStatement("$L.setMaxResults($L)", queryVariableName, queries.result().getLimit().max()); + } + + String pageable = context.getPageableParameterName(); + + if (StringUtils.hasText(pageable)) { + + builder.beginControlFlow("if ($L.isPaged())", pageable); + builder.addStatement("$L.setFirstResult(Long.valueOf($L.getOffset()).intValue())", queryVariableName, pageable); + if (context.returnsSlice() && !context.returnsPage()) { + builder.addStatement("$L.setMaxResults($L.getPageSize() + 1)", queryVariableName, pageable); + } else { + builder.addStatement("$L.setMaxResults($L.getPageSize())", queryVariableName, pageable); + } + builder.endControlFlow(); + } + + return builder.build(); + } + + private CodeBlock createQuery(String queryVariableName, String queryStringNameVariableName, AotQuery query, + MergedAnnotation queryHints) { + + Builder builder = CodeBlock.builder(); + + builder.addStatement("$T $L = this.$L.$L($L)", Query.class, queryVariableName, + context.fieldNameOf(EntityManager.class), query.isNative() ? "createNativeQuery" : "createQuery", + queryStringNameVariableName); + + if (queryHints.isPresent()) { + builder.add(applyHints(queryVariableName, queryHints)); + builder.add("\n"); + } + + for (ParameterBinding binding : query.getParameterBindings()) { + + Object prepare = binding.prepare("s"); + + if (prepare instanceof String prepared && !prepared.equals("s")) { + String format = prepared.replaceAll("%", "%%").replace("s", "%s"); + if (binding.getIdentifier().hasPosition()) { + builder.addStatement("$L.setParameter($L, $S.formatted($L))", queryVariableName, + binding.getIdentifier().getPosition(), format, + context.getParameterNameOfPosition(binding.getIdentifier().getPosition() - 1)); + } else { + builder.addStatement("$L.setParameter($S, $S.formatted($L))", queryVariableName, + binding.getIdentifier().getName(), format, binding.getIdentifier().getName()); + } + } else { + if (binding.getIdentifier().hasPosition()) { + builder.addStatement("$L.setParameter($L, $L)", queryVariableName, binding.getIdentifier().getPosition(), + context.getParameterNameOfPosition(binding.getIdentifier().getPosition() - 1)); + } else { + builder.addStatement("$L.setParameter($S, $L)", queryVariableName, binding.getIdentifier().getName(), + binding.getIdentifier().getName()); + } + } + } + + return builder.build(); + } + + private CodeBlock applyHints(String queryVariableName, MergedAnnotation queryHints) { + + Builder hintsBuilder = CodeBlock.builder(); + MergedAnnotation[] values = queryHints.getAnnotationArray("value", QueryHint.class); + + for (MergedAnnotation hint : values) { + hintsBuilder.addStatement("$L.setHint($S, $S)", queryVariableName, hint.getString("name"), + hint.getString("value")); + } + + return hintsBuilder.build(); + } + + } + + static class QueryExecutionBlockBuilder { + + private final AotRepositoryMethodGenerationContext context; + private String queryVariableName = "query"; + + private QueryExecutionBlockBuilder(AotRepositoryMethodGenerationContext context) { + this.context = context; + } + + public QueryExecutionBlockBuilder referencing(String queryVariableName) { + + this.queryVariableName = queryVariableName; + return this; + } + + public CodeBlock build() { Builder builder = CodeBlock.builder(); @@ -75,7 +290,6 @@ public class JpaCodeBlocks { context.getActualReturnType()); Object actualReturnType = isProjecting ? context.getActualReturnType() : context.getRepositoryInformation().getDomainType(); - builder.add("\n"); if (context.isDeleteMethod()) { @@ -122,171 +336,8 @@ public class JpaCodeBlocks { return builder.build(); } + } - /** - * Builder for the actual query code block. - */ - static class QueryBlockBuilder { - private final AotRepositoryMethodGenerationContext context; - private String queryVariableName = "query"; - private AotQueries queries; - - public QueryBlockBuilder(AotRepositoryMethodGenerationContext context) { - this.context = context; - } - - QueryBlockBuilder usingQueryVariableName(String queryVariableName) { - - this.queryVariableName = queryVariableName; - return this; - } - - QueryBlockBuilder filter(AotQueries query) { - this.queries = query; - return this; - } - - CodeBlock build() { - - boolean isProjecting = context.getActualReturnType() != null - && !ObjectUtils.nullSafeEquals(TypeName.get(context.getRepositoryInformation().getDomainType()), - context.getActualReturnType()); - Object actualReturnType = isProjecting ? context.getActualReturnType() - : context.getRepositoryInformation().getDomainType(); - - CodeBlock.Builder builder = CodeBlock.builder(); - builder.add("\n"); - String queryStringNameVariableName = "%sString".formatted(queryVariableName); - - StringAotQuery query = (StringAotQuery) queries.result(); - builder.addStatement("$T $L = $S", String.class, queryStringNameVariableName, query.getQueryString()); - - String countQueryStringNameVariableName = null; - String countQuyerVariableName = null; - - if (context.returnsPage()) { - - countQueryStringNameVariableName = "count%sString".formatted(StringUtils.capitalize(queryVariableName)); - countQuyerVariableName = "count%s".formatted(StringUtils.capitalize(queryVariableName)); - - StringAotQuery countQuery = (StringAotQuery) queries.count(); - builder.addStatement("$T $L = $S", String.class, countQueryStringNameVariableName, - countQuery.getQueryString()); - } - - // sorting - // TODO: refactor into sort builder - - String sortParameterName = context.getSortParameterName(); - if (sortParameterName == null && context.getPageableParameterName() != null) { - sortParameterName = "%s.getSort()".formatted(context.getPageableParameterName()); - } - - if (StringUtils.hasText(sortParameterName)) { - applySorting(builder, sortParameterName, queryStringNameVariableName, actualReturnType); - } - - addQueryBlock(builder, queryVariableName, queryStringNameVariableName, queries.result()); - - applyLimits(builder); - - if (StringUtils.hasText(countQueryStringNameVariableName)) { - - builder.beginControlFlow("$T $L = () ->", LongSupplier.class, "countAll"); - addQueryBlock(builder, countQuyerVariableName, countQueryStringNameVariableName, queries.count()); - builder.addStatement("return ($T) $L.getSingleResult()", Long.class, countQuyerVariableName); - - // end control flow does not work well with lambdas - builder.unindent(); - builder.add("};\n"); - } - - return builder.build(); - } - - private void applySorting(Builder builder, String sort, String queryString, Object actualReturnType) { - - builder.beginControlFlow("if ($L.isSorted())", sort); - - if (queries.isNative()) { - builder.addStatement("$T declaredQuery = $T.nativeQuery($L)", DeclaredQuery.class, DeclaredQuery.class, - queryString); - } else { - builder.addStatement("$T declaredQuery = $T.jpqlQuery($L)", DeclaredQuery.class, DeclaredQuery.class, - queryString); - } - - builder.addStatement("$L = rewriteQuery(declaredQuery, $L, $T.class)", queryString, sort, actualReturnType); - - builder.endControlFlow(); - } - - private void applyLimits(Builder builder) { - - if (context.isExistsMethod()) { - builder.addStatement("$L.setMaxResults(1)", queryVariableName); - - return; - } - - String limit = context.getLimitParameterName(); - - if (StringUtils.hasText(limit)) { - builder.beginControlFlow("if ($L.isLimited())", limit); - builder.addStatement("$L.setMaxResults($L.max())", queryVariableName, limit); - builder.endControlFlow(); - } else if (queries.result().isLimited()) { - builder.addStatement("$L.setMaxResults($L)", queryVariableName, queries.result().getLimit().max()); - } - - String pageable = context.getPageableParameterName(); - - if (StringUtils.hasText(pageable)) { - - builder.beginControlFlow("if ($L.isPaged())", pageable); - builder.addStatement("$L.setFirstResult(Long.valueOf($L.getOffset()).intValue())", queryVariableName, pageable); - if (context.returnsSlice() && !context.returnsPage()) { - builder.addStatement("$L.setMaxResults($L.getPageSize() + 1)", queryVariableName, pageable); - } else { - builder.addStatement("$L.setMaxResults($L.getPageSize())", queryVariableName, pageable); - } - builder.endControlFlow(); - } - } - - private void addQueryBlock(Builder builder, String queryVariableName, String queryStringNameVariableName, - AotQuery query) { - - builder.addStatement("$T $L = this.$L.$L($L)", Query.class, queryVariableName, - context.fieldNameOf(EntityManager.class), query.isNative() ? "createNativeQuery" : "createQuery", - queryStringNameVariableName); - - for (ParameterBinding binding : query.getParameterBindings()) { - - Object prepare = binding.prepare("s"); - - if (prepare instanceof String prepared && !prepared.equals("s")) { - String format = prepared.replaceAll("%", "%%").replace("s", "%s"); - if (binding.getIdentifier().hasPosition()) { - builder.addStatement("$L.setParameter($L, $S.formatted($L))", queryVariableName, - binding.getIdentifier().getPosition(), format, - context.getParameterNameOfPosition(binding.getIdentifier().getPosition() - 1)); - } else { - builder.addStatement("$L.setParameter($S, $S.formatted($L))", queryVariableName, - binding.getIdentifier().getName(), format, binding.getIdentifier().getName()); - } - } else { - if (binding.getIdentifier().hasPosition()) { - builder.addStatement("$L.setParameter($L, $L)", queryVariableName, binding.getIdentifier().getPosition(), - context.getParameterNameOfPosition(binding.getIdentifier().getPosition() - 1)); - } else { - builder.addStatement("$L.setParameter($S, $L)", queryVariableName, binding.getIdentifier().getName(), - binding.getIdentifier().getName()); - } - } - } - } - } } diff --git a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributor.java b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributor.java index 2d4a92bac..c16f8a6ed 100644 --- a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributor.java +++ b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributor.java @@ -125,8 +125,8 @@ public class JpaRepositoryContributor extends RepositoryContributor { aotQueries = buildPartTreeQuery(context, query); } - body.addCode(JpaCodeBlocks.queryBlockBuilder(context).filter(aotQueries).build()); - body.addCode(JpaCodeBlocks.queryExecutionBlockBuilder(context).build()); + body.addCode(JpaCodeBlocks.queryBuilder(context).filter(aotQueries).queryHints(queryHints).build()); + body.addCode(JpaCodeBlocks.executionBuilder(context).build()); }); } diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributorIntegrationTests.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributorIntegrationTests.java index 582b47627..bfa807708 100644 --- a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributorIntegrationTests.java +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/JpaRepositoryContributorIntegrationTests.java @@ -297,6 +297,12 @@ class JpaRepositoryContributorIntegrationTests { "kylo@new-empire.com", "luke@jedi.org", "vader@empire.com"); } + @Test + void shouldApplyQueryHints() { + assertThatIllegalArgumentException().isThrownBy(() -> fragment.findHintedByLastname("Skywalker")) + .withMessageContaining("No enum constant jakarta.persistence.CacheStoreMode.foo"); + } + @Test void testDerivedFinderReturningPageOfProjections() { @@ -344,8 +350,9 @@ class JpaRepositoryContributorIntegrationTests { // interface projections // named queries + // dynamic projections + // class type parameter - // query hints // entity graphs // native queries // delete diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/UserRepository.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/UserRepository.java index 4e8088fa3..d783c9cdf 100644 --- a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/UserRepository.java +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/aot/generated/UserRepository.java @@ -15,6 +15,8 @@ */ package org.springframework.data.jpa.repository.aot.generated; +import jakarta.persistence.QueryHint; + import java.util.List; import java.util.Optional; @@ -26,6 +28,7 @@ import org.springframework.data.domain.Sort; import org.springframework.data.jpa.domain.sample.User; import org.springframework.data.jpa.repository.Modifying; import org.springframework.data.jpa.repository.Query; +import org.springframework.data.jpa.repository.QueryHints; import org.springframework.data.repository.CrudRepository; /** @@ -128,6 +131,9 @@ public interface UserRepository extends CrudRepository { List findByLastname(String lastname); + @QueryHints(value = { @QueryHint(name = "jakarta.persistence.cache.storeMode", value = "foo") }, forCounting = false) + List findHintedByLastname(String lastname); + List findByLastnameStartingWithOrderByFirstname(String lastname, Limit limit); List findByLastname(String lastname, Sort sort);