diff --git a/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java b/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java index 3a173b5a2..92d1b8059 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java @@ -62,7 +62,9 @@ public abstract class QueryUtils { public static final String DELETE_ALL_QUERY_STRING = "delete from %s x"; private static final String DEFAULT_ALIAS = "x"; - private static final String COUNT_REPLACEMENT = "select count($3$5) $4$5$6"; + private static final String COUNT_REPLACEMENT_TEMPLATE = "select count(%s) $5$6$7"; + private static final String SIMPLE_COUNT_VALUE = "$2"; + private static final String COMPLEX_COUNT_VALUE = "$3$6"; private static final Pattern ALIAS_MATCH; private static final Pattern COUNT_MATCH; @@ -90,7 +92,7 @@ public abstract class QueryUtils { ALIAS_MATCH = compile(builder.toString(), CASE_INSENSITIVE); builder = new StringBuilder(); - builder.append("(select\\s+((distinct )?.+?)\\s+)?(from\\s+"); + builder.append("(select\\s+((distinct )?(.+)?)\\s+)?(from\\s+"); builder.append(IDENTIFIER); builder.append("(?:\\s+as)?\\s+)"); builder.append(IDENTIFIER_GROUP); @@ -326,7 +328,12 @@ public abstract class QueryUtils { Assert.hasText(originalQuery); Matcher matcher = COUNT_MATCH.matcher(originalQuery); - return matcher.replaceFirst(COUNT_REPLACEMENT); + String variable = matcher.matches() ? matcher.group(4) : null; + boolean useVariable = StringUtils.hasText(variable) && !variable.startsWith("new") + && !variable.startsWith("count("); + + return matcher.replaceFirst(String.format(COUNT_REPLACEMENT_TEMPLATE, useVariable ? SIMPLE_COUNT_VALUE + : COMPLEX_COUNT_VALUE)); } /** diff --git a/src/main/java/org/springframework/data/jpa/repository/query/SimpleJpaQuery.java b/src/main/java/org/springframework/data/jpa/repository/query/SimpleJpaQuery.java index 8423541a1..9c43d281b 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/SimpleJpaQuery.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/SimpleJpaQuery.java @@ -51,8 +51,6 @@ final class SimpleJpaQuery extends AbstractJpaQuery { this.method = method; this.query = new StringQuery(queryString); - this.countQuery = new StringQuery(method.getCountQuery() == null ? QueryUtils.createCountQueryFor(queryString) - : method.getCountQuery()); Parameters parameters = method.getParameters(); boolean hasPagingOrSortingParameter = parameters.hasPageableParameter() || parameters.hasSortParameter(); @@ -61,15 +59,28 @@ final class SimpleJpaQuery extends AbstractJpaQuery { throw new IllegalStateException("Cannot use native queries with dynamic sorting and/or pagination!"); } - // Try to create a Query object already to fail fast + String preparedQueryString = this.query.getQuery(); + if (!method.isNativeQuery()) { - try { - em.createQuery(query.getQuery()); - } catch (RuntimeException e) { - // Needed as there's ambiguities in how an invalid query string shall be expressed by the persistence provider - // http://java.net/projects/jpa-spec/lists/jsr338-experts/archive/2012-07/message/17 - throw e instanceof IllegalArgumentException ? e : new IllegalArgumentException(e); - } + validateQuery(preparedQueryString, em); + } + + this.countQuery = new StringQuery(method.getCountQuery() != null ? method.getCountQuery() + : QueryUtils.createCountQueryFor(preparedQueryString)); + + if (!method.isNativeQuery()) { + validateQuery(this.countQuery.getQuery(), em); + } + } + + private final void validateQuery(String query, EntityManager em) { + + try { + em.createQuery(query); + } catch (RuntimeException e) { + // Needed as there's ambiguities in how an invalid query string shall be expressed by the persistence provider + // http://java.net/projects/jpa-spec/lists/jsr338-experts/archive/2012-07/message/17 + throw e instanceof IllegalArgumentException ? e : new IllegalArgumentException(e); } } diff --git a/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java b/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java index ce627f308..522556181 100644 --- a/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java +++ b/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java @@ -203,6 +203,16 @@ public class QueryUtilsUnitTests { assertThat(applySorting(query, sort, "p"), endsWith("order by p.lastname asc, lower(p.firstname) asc")); } + /** + * @see DATAJPA-342 + */ + @Test + public void usesReturnedVariableInCOuntProjectionIfSet() { + + assertCountQuery("select distinct m.genre from Media m where m.user = ?1 order by m.genre asc", + "select count(distinct m.genre) from Media m where m.user = ?1 order by m.genre asc"); + } + private void assertCountQuery(String originalQuery, String countQuery) { assertThat(createCountQueryFor(originalQuery), is(countQuery)); }