diff --git a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java index da041122d..56151a04b 100644 --- a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java +++ b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/QueryUtils.java @@ -525,14 +525,19 @@ public abstract class QueryUtils { boolean useVariable = StringUtils.hasText(variable) // && !variable.startsWith(" new") // && !variable.startsWith("count(") // - && !variable.contains(",") // - && !variable.contains("*"); + && !variable.contains(","); String complexCountValue = matcher.matches() && StringUtils.hasText(matcher.group(COMPLEX_COUNT_FIRST_INDEX)) ? COMPLEX_COUNT_VALUE : COMPLEX_COUNT_LAST_VALUE; String replacement = useVariable ? SIMPLE_COUNT_VALUE : complexCountValue; + + String alias = QueryUtils.detectAlias(originalQuery); + if("*".equals(variable) && alias != null) { + replacement = alias; + } + countQuery = matcher.replaceFirst(String.format(COUNT_REPLACEMENT_TEMPLATE, replacement)); } else { countQuery = matcher.replaceFirst(String.format(COUNT_REPLACEMENT_TEMPLATE, countProjection)); diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java index e40535c1e..1b56f3629 100644 --- a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java @@ -712,6 +712,26 @@ class QueryEnhancerUnitTests { assertThat(result).containsIgnoringCase("order by dd.institutesIds"); } + + @Test //GH-2511 + void countQueryUsesCorrectVariable() { + StringQuery nativeQuery = new StringQuery("SELECT * FROM User WHERE created_at > $1", true); + QueryEnhancer queryEnhancer = getEnhancer(nativeQuery); + String countQueryFor = queryEnhancer.createCountQueryFor(); + assertThat(countQueryFor).isEqualTo("SELECT count(*) FROM User WHERE created_at > $1"); + + nativeQuery = new StringQuery("SELECT * FROM (select * from test) ",true); + queryEnhancer = getEnhancer(nativeQuery); + countQueryFor = queryEnhancer.createCountQueryFor(); + assertThat(countQueryFor).isEqualTo("SELECT count(*) FROM (SELECT * FROM test)"); + + nativeQuery = new StringQuery("SELECT * FROM (select * from test) as test",true); + queryEnhancer = getEnhancer(nativeQuery); + countQueryFor = queryEnhancer.createCountQueryFor(); + assertThat(countQueryFor).isEqualTo("SELECT count(test) FROM (SELECT * FROM test) AS test"); + } + + public static Stream detectsJoinAliasesCorrectlySource() { return Stream.of( // diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java index 9a3e86965..e6a9a3f4b 100644 --- a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/QueryUtilsUnitTests.java @@ -638,4 +638,21 @@ class QueryUtilsUnitTests { "select * from (select * from user order by 1, 2, 3 desc limit 10) u order by u.active asc, age desc"); } + @Test //GH-2511 + void countQueryUsesCorrectVariable() { + String countQueryFor = createCountQueryFor("SELECT * FROM User WHERE created_at > $1"); + assertThat(countQueryFor).isEqualTo("select count(*) FROM User WHERE created_at > $1"); + + countQueryFor = createCountQueryFor("SELECT * FROM mytable WHERE nr = :number AND kon = :kon AND datum >= '2019-01-01'"); + assertThat(countQueryFor).isEqualTo("select count(*) FROM mytable WHERE nr = :number AND kon = :kon AND datum >= '2019-01-01'"); + + countQueryFor = createCountQueryFor("SELECT * FROM context ORDER BY time"); + assertThat(countQueryFor).isEqualTo("select count(*) FROM context"); + + countQueryFor = createCountQueryFor("select * FROM users_statuses WHERE (user_created_at BETWEEN $1 AND $2)"); + assertThat(countQueryFor).isEqualTo("select count(*) FROM users_statuses WHERE (user_created_at BETWEEN $1 AND $2)"); + + countQueryFor = createCountQueryFor("SELECT * FROM users_statuses us WHERE (user_created_at BETWEEN :fromDate AND :toDate)"); + assertThat(countQueryFor).isEqualTo("select count(us) FROM users_statuses us WHERE (user_created_at BETWEEN :fromDate AND :toDate)"); + } }