diff --git a/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java b/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java index 3c4d038c0..d7aad1e4c 100644 --- a/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java +++ b/src/main/java/org/springframework/data/jpa/repository/query/JSqlParserQueryEnhancer.java @@ -24,11 +24,14 @@ import net.sf.jsqlparser.expression.Expression; import net.sf.jsqlparser.expression.Function; import net.sf.jsqlparser.parser.CCJSqlParserUtil; import net.sf.jsqlparser.schema.Column; +import net.sf.jsqlparser.statement.Statement; +import net.sf.jsqlparser.statement.delete.Delete; import net.sf.jsqlparser.statement.select.OrderByElement; import net.sf.jsqlparser.statement.select.PlainSelect; import net.sf.jsqlparser.statement.select.Select; import net.sf.jsqlparser.statement.select.SelectExpressionItem; import net.sf.jsqlparser.statement.select.SelectItem; +import net.sf.jsqlparser.statement.update.Update; import java.util.ArrayList; import java.util.Collections; @@ -54,20 +57,49 @@ import org.springframework.util.StringUtils; public class JSqlParserQueryEnhancer implements QueryEnhancer { private final DeclaredQuery query; + private final ParsedType parsedType; /** * @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()); + + if (statement instanceof Update) { + return ParsedType.UPDATE; + } else if (statement instanceof Delete) { + return ParsedType.DELETE; + } else if (statement instanceof Select) { + return ParsedType.SELECT; + } else { + return ParsedType.SELECT; + } + + } catch (JSQLParserException e) { + throw new IllegalArgumentException("The query you provided is not a valid SQL Query!", e); + } } @Override public String applySorting(Sort sort, @Nullable String alias) { - String queryString = query.getQueryString(); Assert.hasText(queryString, "Query must not be null or empty!"); + if (this.parsedType != ParsedType.SELECT) { + return queryString; + } + if (sort.isUnsorted()) { return queryString; } @@ -120,6 +152,10 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { */ Set getSelectionAliases() { + if (this.parsedType != ParsedType.SELECT) { + return new HashSet<>(); + } + Select selectStatement = parseSelectStatement(this.query.getQueryString()); PlainSelect selectBody = (PlainSelect) selectStatement.getSelectBody(); return this.getSelectionAliases(selectBody); @@ -132,6 +168,9 @@ 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()); } @@ -211,6 +250,10 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { @Nullable private String detectAlias(String query) { + if (this.parsedType != ParsedType.SELECT) { + return null; + } + Select selectStatement = parseSelectStatement(query); PlainSelect selectBody = (PlainSelect) selectStatement.getSelectBody(); return detectAlias(selectBody); @@ -233,6 +276,10 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { @Override public String createCountQueryFor(@Nullable String countProjection) { + if (this.parsedType != ParsedType.SELECT) { + return this.query.getQueryString(); + } + Assert.hasText(this.query.getQueryString(), "OriginalQuery must not be null or empty!"); Select selectStatement = parseSelectStatement(this.query.getQueryString()); @@ -278,6 +325,10 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { @Override public String getProjection() { + if (this.parsedType != ParsedType.SELECT) { + return ""; + } + Assert.hasText(query.getQueryString(), "Query must not be null or empty!"); Select selectStatement = parseSelectStatement(query.getQueryString()); @@ -327,3 +378,15 @@ public class JSqlParserQueryEnhancer implements QueryEnhancer { return this.query; } } + +/** + * An enum to represent the top level parsed statement of the provided query. + * + */ +enum ParsedType { + DELETE, UPDATE, SELECT; +} diff --git a/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java b/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java index c44584efb..65cbf3395 100644 --- a/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java +++ b/src/test/java/org/springframework/data/jpa/repository/UserRepositoryTests.java @@ -2678,6 +2678,19 @@ public class UserRepositoryTests { assertThat(repository.exists(hundredYearsOld)).isTrue(); } + @Test // GH-2555 + void modifyingUpdateNativeQueryWorksWithJSQLParser() { + flushTestUsers(); + + Optional byIdUser = repository.findById(firstUser.getId()); + assertThat(byIdUser).isPresent().map(User::isActive).get().isEqualTo(true); + + repository.setActiveToFalseWithModifyingNative(byIdUser.get().getId()); + + Optional afterUpdate = repository.findById(firstUser.getId()); + assertThat(afterUpdate).isPresent().map(User::isActive).get().isEqualTo(false); + } + @Test // GH-2045, GH-425 public void correctlyBuildSortClauseWhenSortingByFunctionAliasAndFunctionContainsPositionalParameters() { repository.findAllAndSortByFunctionResultPositionalParameter("prefix", "suffix", Sort.by("idWithPrefixAndSuffix")); diff --git a/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java b/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java index c181425f1..5b8487cfb 100644 --- a/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java +++ b/src/test/java/org/springframework/data/jpa/repository/query/QueryEnhancerUnitTests.java @@ -731,6 +731,25 @@ class QueryEnhancerUnitTests { assertThat(countQueryFor).isEqualTo("SELECT count(test) FROM (SELECT * FROM test) AS test"); } + @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); + String projectionNotConsideringQueryType = QueryUtils.getProjection(modifyingQuery); + boolean constructorExpressionNotConsideringQueryType = QueryUtils.hasConstructorExpression(modifyingQuery); + String countQueryForNotConsiderQueryType = QueryUtils.createCountQueryFor(modifyingQuery); + + StringQuery modiQuery = new StringQuery(modifyingQuery, true); + + assertThat(modiQuery.getAlias()).isEqualToIgnoringCase(aliasNotConsideringQueryType); + assertThat(modiQuery.getProjection()).isEqualToIgnoringCase(projectionNotConsideringQueryType); + assertThat(modiQuery.hasConstructorExpression()).isEqualTo(constructorExpressionNotConsideringQueryType); + + assertThat(countQueryForNotConsiderQueryType).isEqualToIgnoringCase(modifyingQuery); + assertThat(QueryEnhancerFactory.forQuery(modiQuery).createCountQueryFor()).isEqualToIgnoringCase(modifyingQuery); + } + public static Stream detectsJoinAliasesCorrectlySource() { return Stream.of( // diff --git a/src/test/java/org/springframework/data/jpa/repository/sample/UserRepository.java b/src/test/java/org/springframework/data/jpa/repository/sample/UserRepository.java index 32ca93f16..20d69f931 100644 --- a/src/test/java/org/springframework/data/jpa/repository/sample/UserRepository.java +++ b/src/test/java/org/springframework/data/jpa/repository/sample/UserRepository.java @@ -637,6 +637,11 @@ public interface UserRepository List findAllAndSortByFunctionResultNamedParameter(@Param("namedParameter1") String namedParameter1, @Param("namedParameter2") String namedParameter2, Sort sort); + // GH-2555 + @Modifying(clearAutomatically = true) + @Query(value = "update SD_User u set u.active = false where u.id = :userId", nativeQuery = true) + void setActiveToFalseWithModifyingNative(@Param("userId") int userId); + interface RolesAndFirstname { String getFirstname();