From cc1eb9ce2b512a3928e55bd190397b751814cd9b Mon Sep 17 00:00:00 2001 From: "Greg L. Turnquist" Date: Fri, 3 Jun 2022 14:44:32 -0500 Subject: [PATCH] Polishing. See #2555. --- .../query/JSqlParserQueryEnhancer.java | 31 +++++++++++-------- .../jpa/repository/UserRepositoryTests.java | 1 + .../query/QueryEnhancerUnitTests.java | 1 + 3 files changed, 20 insertions(+), 13 deletions(-) diff --git a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java index d7aad1e4c..d8322ec8e 100644 --- a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java +++ b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java @@ -63,16 +63,18 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { * @param query the query we want to enhance. Must not be {@literal null}. */ public JSqlParserQueryEnhancer(DeclaredQuery query) { + this.query = query; this.parsedType = detectParsedType(); } /** * Detects what type of query is provided. - * + * * @return the parsed type */ private ParsedType detectParsedType() { + try { Statement statement = CCJSqlParserUtil.parse(this.query.getQueryString()); @@ -85,7 +87,6 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { } else { return ParsedType.SELECT; } - } catch (JSQLParserException e) { throw new IllegalArgumentException("The query you provided is not a valid SQL Query!", e); } @@ -93,6 +94,7 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { @Override public String applySorting(Sort sort, @Nullable String alias) { + String queryString = query.getQueryString(); Assert.hasText(queryString, "Query must not be null or empty!"); @@ -168,9 +170,11 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { * @return a {@literal Set} of aliases used in the query. Guaranteed to be not {@literal null}. */ private Set getJoinAliases(String query) { + if (this.parsedType != ParsedType.SELECT) { return new HashSet<>(); } + return getJoinAliases((PlainSelect) parseSelectStatement(query).getSelectBody()); } @@ -377,16 +381,17 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { public DeclaredQuery getQuery() { return this.query; } -} -/** - * An enum to represent the top level parsed statement of the provided query. - * - */ -enum ParsedType { - DELETE, UPDATE, SELECT; + /** + * An enum to represent the top level parsed statement of the provided query. + * + */ + enum ParsedType { + DELETE, UPDATE, SELECT; + } + } diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java index 4b7bb6b5f..7bc48e30a 100644 --- a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java @@ -2867,6 +2867,7 @@ public class UserRepositoryTests { @Test // GH-2555 void modifyingUpdateNativeQueryWorksWithJSQLParser() { + flushTestUsers(); Optional byIdUser = repository.findById(firstUser.getId()); 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 5b8487cfb..3572fa403 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 @@ -733,6 +733,7 @@ class QueryEnhancerUnitTests { @Test // GH-2555 void modifyingQueriesAreDetectedCorrectly() { + String modifyingQuery = "update userinfo user set user.is_in_treatment = false where user.id = :userId"; String aliasNotConsideringQueryType = QueryUtils.detectAlias(modifyingQuery);