Add query hint support.

See #3830
This commit is contained in:
Mark Paluch
2025-03-26 08:56:51 +01:00
parent c399ca2b54
commit 802a8db1d1
4 changed files with 241 additions and 177 deletions

View File

@@ -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> 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> 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> 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> queryHints) {
Builder hintsBuilder = CodeBlock.builder();
MergedAnnotation<QueryHint>[] values = queryHints.getAnnotationArray("value", QueryHint.class);
for (MergedAnnotation<QueryHint> 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());
}
}
}
}
}
}

View File

@@ -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());
});
}

View File

@@ -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

View File

@@ -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<User, Integer> {
List<User> findByLastname(String lastname);
@QueryHints(value = { @QueryHint(name = "jakarta.persistence.cache.storeMode", value = "foo") }, forCounting = false)
List<User> findHintedByLastname(String lastname);
List<User> findByLastnameStartingWithOrderByFirstname(String lastname, Limit limit);
List<User> findByLastname(String lastname, Sort sort);