From 7bcc8042c1e5c6871b8d91cd9dda9525af9911b7 Mon Sep 17 00:00:00 2001 From: "Greg L. Turnquist" Date: Wed, 4 May 2022 11:31:16 -0500 Subject: [PATCH] Polishing. See #2260 (c93aa25), #2500, #2518. --- .../data/jpa/repository/query/QueryUtils.java | 44 ++++++++++-------- .../repository/query/QueryUtilsUnitTests.java | 45 +++++++++++++------ 2 files changed, 57 insertions(+), 32 deletions(-) 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 30cb53f6a..fb7a0a0bb 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 @@ -433,6 +433,7 @@ public abstract class QueryUtils { @Nullable @Deprecated public static String detectAlias(String query) { + String alias = null; Matcher matcher = ALIAS_MATCH.matcher(removeSubqueries(query)); while (matcher.find()) { @@ -442,23 +443,25 @@ public abstract class QueryUtils { } /** - * Remove subqueries from the query, in order to identify the correct alias - * in order by clauses. If the entire query is surrounded by parenthesis, the - * outermost parenthesis are not removed. + * Remove subqueries from the query, in order to identify the correct alias in order by clauses. If the entire query + * is surrounded by parenthesis, the outermost parenthesis are not removed. * * @param query * @return query with all subqueries removed. */ static String removeSubqueries(String query) { + if (!StringUtils.hasText(query)) { return query; } - final List opens = new ArrayList<>(); - final List closes = new ArrayList<>(); - final List closeMatches = new ArrayList<>(); - for (int i=0; i opens = new ArrayList<>(); + List closes = new ArrayList<>(); + List closeMatches = new ArrayList<>(); + + for (int i = 0; i < query.length(); i++) { + + char c = query.charAt(i); if (c == '(') { opens.add(i); } else if (c == ')') { @@ -467,18 +470,19 @@ public abstract class QueryUtils { } } - final StringBuilder sb = new StringBuilder(query); - final boolean startsWithParen = STARTS_WITH_PAREN.matcher(query).find(); - for (int i=opens.size()-1; i>=(startsWithParen?1:0); i--) { - final Integer open = opens.get(i); - final Integer close = findClose(open, closes, closeMatches) + 1; + StringBuilder sb = new StringBuilder(query); + boolean startsWithParen = STARTS_WITH_PAREN.matcher(query).find(); + for (int i = opens.size() - 1; i >= (startsWithParen ? 1 : 0); i--) { + Integer open = opens.get(i); + Integer close = findClose(open, closes, closeMatches) + 1; if (close > open) { - final String subquery = sb.substring(open, close); - final Matcher matcher = PARENS_TO_REMOVE.matcher(subquery); + + String subquery = sb.substring(open, close); + Matcher matcher = PARENS_TO_REMOVE.matcher(subquery); if (matcher.find()) { - sb.replace(open, close, new String(new char[close-open]).replace('\0', ' ')); + sb.replace(open, close, new String(new char[close - open]).replace('\0', ' ')); } } } @@ -487,8 +491,10 @@ public abstract class QueryUtils { } private static Integer findClose(final Integer open, final List closes, final List closeMatches) { - for (int i=0; i open && !closeMatches.get(i)) { closeMatches.set(i, Boolean.TRUE); return close; @@ -594,7 +600,7 @@ public abstract class QueryUtils { String replacement = useVariable ? SIMPLE_COUNT_VALUE : complexCountValue; String alias = QueryUtils.detectAlias(originalQuery); - if("*".equals(variable) && alias != null) { + if ("*".equals(variable) && alias != null) { replacement = alias; } 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 4e8e244b6..08c182083 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 @@ -120,13 +120,19 @@ class QueryUtilsUnitTests { assertThat(detectAlias("select u from T05User u")).isEqualTo("u"); assertThat(detectAlias("select u from User u where not exists (from User u2)")).isEqualTo("u"); assertThat(detectAlias("(select u from User u where not exists (from User u2))")).isEqualTo("u"); - assertThat(detectAlias("(select u from User u where not exists ((from User u2 where not exists (from User u3))))")).isEqualTo("u"); - assertThat(detectAlias("from Foo f left join f.bar b with type(b) = BarChild where (f.id = (select max(f.id) from Foo f2 where type(f2) = FooChild) or 1 <> 1) and 1=1")).isEqualTo("f"); - assertThat(detectAlias("(from Foo f max(f) ((((select * from Foo f2 (from Foo f3) max(*)) (from Foo f4)) max(f5)) (f6)) (from Foo f7))")).isEqualTo("f"); + assertThat(detectAlias("(select u from User u where not exists ((from User u2 where not exists (from User u3))))")) + .isEqualTo("u"); + assertThat(detectAlias( + "from Foo f left join f.bar b with type(b) = BarChild where (f.id = (select max(f.id) from Foo f2 where type(f2) = FooChild) or 1 <> 1) and 1=1")) + .isEqualTo("f"); + assertThat(detectAlias( + "(from Foo f max(f) ((((select * from Foo f2 (from Foo f3) max(*)) (from Foo f4)) max(f5)) (f6)) (from Foo f7))")) + .isEqualTo("f"); } @Test // GH-2260 void testRemoveSubqueries() throws Exception { + // boundary conditions assertThat(removeSubqueries(null)).isNull(); assertThat(removeSubqueries("")).isEmpty(); @@ -145,11 +151,19 @@ class QueryUtilsUnitTests { assertThat(removeSubqueries("select u from User u")).isEqualTo("select u from User u"); assertThat(removeSubqueries("select u from com.acme.User u")).isEqualTo("select u from com.acme.User u"); assertThat(removeSubqueries("select u from T05User u")).isEqualTo("select u from T05User u"); - assertThat(normalizeWhitespace(removeSubqueries("select u from User u where not exists (from User u2)"))).isEqualTo("select u from User u where not exists"); - assertThat(normalizeWhitespace(removeSubqueries("(select u from User u where not exists (from User u2))"))).isEqualTo("(select u from User u where not exists )"); - assertThat(normalizeWhitespace(removeSubqueries("select u from User u where not exists (from User u2 where not exists (from User u3))"))).isEqualTo("select u from User u where not exists"); - assertThat(normalizeWhitespace(removeSubqueries("select u from User u where not exists ((from User u2 where not exists (from User u3)))"))).isEqualTo("select u from User u where not exists ( )"); - assertThat(normalizeWhitespace(removeSubqueries("(select u from User u where not exists ((from User u2 where not exists (from User u3))))"))).isEqualTo("(select u from User u where not exists ( ))"); + assertThat(normalizeWhitespace(removeSubqueries("select u from User u where not exists (from User u2)"))) + .isEqualTo("select u from User u where not exists"); + assertThat(normalizeWhitespace(removeSubqueries("(select u from User u where not exists (from User u2))"))) + .isEqualTo("(select u from User u where not exists )"); + assertThat(normalizeWhitespace( + removeSubqueries("select u from User u where not exists (from User u2 where not exists (from User u3))"))) + .isEqualTo("select u from User u where not exists"); + assertThat(normalizeWhitespace( + removeSubqueries("select u from User u where not exists ((from User u2 where not exists (from User u3)))"))) + .isEqualTo("select u from User u where not exists ( )"); + assertThat(normalizeWhitespace( + removeSubqueries("(select u from User u where not exists ((from User u2 where not exists (from User u3))))"))) + .isEqualTo("(select u from User u where not exists ( ))"); } private String normalizeWhitespace(String s) { @@ -690,16 +704,21 @@ class QueryUtilsUnitTests { 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 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)"); + 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)"); + 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)"); } }