diff --git a/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitor.java b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitor.java new file mode 100644 index 000000000..ade370c0f --- /dev/null +++ b/spring-data-jpa/src/main/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitor.java @@ -0,0 +1,740 @@ +/* + * Copyright 2025 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.jpa.repository.query; + +import jakarta.persistence.criteria.CriteriaBuilder; +import jakarta.persistence.criteria.Expression; +import jakarta.persistence.criteria.From; +import jakarta.persistence.criteria.LocalDateTimeField; +import jakarta.persistence.criteria.Path; +import jakarta.persistence.criteria.TemporalField; + +import java.math.BigDecimal; +import java.math.BigInteger; +import java.time.temporal.Temporal; +import java.util.Collection; +import java.util.HexFormat; + +import org.antlr.v4.runtime.CharStreams; +import org.antlr.v4.runtime.CommonTokenStream; +import org.antlr.v4.runtime.tree.ParseTree; +import org.antlr.v4.runtime.tree.TerminalNode; +import org.hibernate.query.criteria.HibernateCriteriaBuilder; + +import org.springframework.data.domain.Sort; +import org.springframework.data.jpa.domain.JpaSort; +import org.springframework.data.mapping.PropertyPath; +import org.springframework.util.Assert; + +/** + * Parses the content of {@link JpaSort#unsafe(String...)} as an HQL {@literal sortExpression} and renders that into a + * JPA Criteria {@link Expression}. + * + * @author Greg Turnquist + * @author Mark Paluch + * @since 4.0 + */ +@SuppressWarnings("ConstantValue") +class HqlOrderExpressionVisitor extends HqlBaseVisitor> { + + private final CriteriaBuilder cb; + private final Path from; + private static String UNSUPPORTED_TEMPLATE = "We can't handle %s in an ORDER BY clause through JpaSort.unsafe"; + + HqlOrderExpressionVisitor(CriteriaBuilder cb, Path from) { + this.cb = cb; + this.from = from; + } + + /** + * Extract the {@link org.springframework.data.jpa.domain.JpaSort.JpaOrder}'s property and parse it as an HQL + * {@literal sortExpression}. + * + * @param jpaOrder + * @return criteriaExpression + */ + Expression createCriteriaExpression(Sort.Order jpaOrder) { + + String orderByProperty = jpaOrder.getProperty(); + HqlLexer lexer = new HqlLexer(CharStreams.fromString(orderByProperty)); + HqlParser parser = new HqlParser(new CommonTokenStream(lexer)); + + JpaQueryEnhancer.configureParser(orderByProperty, "ORDER BY expression", lexer, parser); + + HqlParser.SortExpressionContext ctx = parser.sortExpression(); + + if (ctx == null) { + throw new IllegalArgumentException("No sort expression provided"); + } + + return visitRequired(ctx); + } + + @Override + public Expression visitSortExpression(HqlParser.SortExpressionContext ctx) { + + if (ctx.identifier() != null) { + HqlParser.IdentifierContext identifier = ctx.identifier(); + + return from.get(getString(identifier)); + } else if (ctx.INTEGER_LITERAL() != null) { + return cb.literal(Integer.valueOf(ctx.INTEGER_LITERAL().getText())); + } else if (ctx.expression() != null) { + return visitRequired(ctx.expression()); + } else { + return null; + } + } + + @Override + @SuppressWarnings("rawtypes") + public Expression visitRelationalExpression(HqlParser.RelationalExpressionContext ctx) { + + Expression left = visitRequired(ctx.expression(0)); + Expression right = visitRequired(ctx.expression(1)); + String op = ctx.op.getText(); + + if (op.equals("=")) { + return cb.equal(left, right); + } else if (op.equals(">")) { + return cb.greaterThan(left, right); + } else if (op.equals(">=")) { + return cb.greaterThanOrEqualTo(left, right); + } else if (op.equals("<")) { + return cb.lessThan(left, right); + } else if (op.equals("<=")) { + return cb.lessThanOrEqualTo(left, right); + } else if (op.equals("<>") || op.equals("!=") || op.equals("^=")) { + return cb.notEqual(left, right); + } else { + throw new UnsupportedOperationException("Unsupported comparison operator: " + op); + } + } + + @Override + @SuppressWarnings("rawtypes") + public Expression visitBetweenExpression(HqlParser.BetweenExpressionContext ctx) { + + Expression condition = visitRequired(ctx.expression(0)); + Expression lower = visitRequired(ctx.expression(1)); + Expression upper = visitRequired(ctx.expression(2)); + + if (ctx.NOT() == null) { + return cb.between(condition, lower, upper); + } else { + return cb.between(condition, lower, upper).not(); + } + } + + @SuppressWarnings("unchecked") + @Override + public Expression visitIsBooleanPredicate(HqlParser.IsBooleanPredicateContext ctx) { + + Expression condition = visitRequired(ctx.expression()); + + if (ctx.NULL() != null) { + if (ctx.NOT() == null) { + return cb.isNull(condition); + } else { + return cb.isNotNull(condition); + } + } + + if (ctx.EMPTY() != null) { + if (ctx.NOT() == null) { + return cb.isEmpty((Expression>) condition); + } else { + return cb.isNotEmpty((Expression>) condition); + } + } + + if (ctx.TRUE() != null) { + if (ctx.NOT() == null) { + return cb.isTrue((Expression) condition); + } else { + return cb.isFalse((Expression) condition); + } + } + + if (ctx.FALSE() != null) { + if (ctx.NOT() == null) { + return cb.isFalse((Expression) condition); + } else { + return cb.isTrue((Expression) condition); + } + } + + return null; + } + + @Override + public Expression visitStringPatternMatching(HqlParser.StringPatternMatchingContext ctx) { + Expression condition = visitRequired(ctx.expression(0)); + Expression match = visitRequired(ctx.expression(1)); + Expression escape = ctx.ESCAPE() != null ? charLiteralOf(ctx.ESCAPE()) : null; + + if (ctx.LIKE() != null) { + if (ctx.NOT() == null) { + return escape == null // + ? cb.like(condition, match) // + : cb.like(condition, match, escape); + } else { + return escape == null // + ? cb.notLike(condition, match) // + : cb.notLike(condition, match, escape); + } + } else { + HibernateCriteriaBuilder hcb = (HibernateCriteriaBuilder) cb; + if (ctx.NOT() == null) { + return escape == null // + ? hcb.ilike(condition, match) // + : hcb.ilike(condition, match, escape); + } else { + return escape == null // + ? hcb.notIlike(condition, match) // + : hcb.notIlike(condition, match, escape); + } + } + } + + @Override + public Expression visitInExpression(HqlParser.InExpressionContext ctx) { + + if (ctx.inList().simplePath() != null) { + throw new UnsupportedOperationException( + String.format(UNSUPPORTED_TEMPLATE, "IN clause with ELEMENTS or INDICES argument")); + } else if (ctx.inList().subquery() != null) { + throw new UnsupportedOperationException(String.format(UNSUPPORTED_TEMPLATE, "IN clause with a subquery")); + } else if (ctx.inList().parameter() != null) { + throw new UnsupportedOperationException(String.format(UNSUPPORTED_TEMPLATE, "IN clause with a parameter")); + } + + CriteriaBuilder.In in = cb.in(visit(ctx.expression())); + + ctx.inList().expressionOrPredicate() + .forEach(expressionOrPredicateContext -> in.value(visit(expressionOrPredicateContext))); + + if (ctx.NOT() == null) { + return in; + } + return in.not(); + + } + + @Override + public Expression visitGenericFunction(HqlParser.GenericFunctionContext ctx) { + + String functionName = ctx.genericFunctionName().getText(); + + if (ctx.genericFunctionArguments() == null) { + return cb.function(functionName, Object.class); + } + + Expression[] arguments = ctx.genericFunctionArguments().expressionOrPredicate().stream() // + .map(expressionOrPredicateContext -> visitRequired(expressionOrPredicateContext)) // + .toArray(Expression[]::new); + return cb.function(functionName, Object.class, arguments); + + } + + @Override + public Expression visitCastFunction(HqlParser.CastFunctionContext ctx) { + throw new UnsupportedOperationException("Sorting using CAST ist not supported"); + } + + @Override + public Expression visitTreatedNavigablePath(HqlParser.TreatedNavigablePathContext ctx) { + throw new UnsupportedOperationException("Sorting using TREAT ist not supported"); + } + + @Override + @SuppressWarnings({ "rawtypes", "unchecked" }) + public Expression visitExtractFunction(HqlParser.ExtractFunctionContext ctx) { + + Expression expr = visitRequired(ctx.expression()); + TemporalField temporalField = ctx.extractField() != null ? getTemporalField(ctx.extractField()) + : getTemporalField(ctx.datetimeField()); + + return cb.extract(temporalField, expr); + } + + private TemporalField getTemporalField(HqlParser.DatetimeFieldContext ctx) { + + if (ctx.YEAR() != null) { + return LocalDateTimeField.YEAR; + } + + if (ctx.MONTH() != null) { + return LocalDateTimeField.MONTH; + } + + if (ctx.QUARTER() != null) { + return LocalDateTimeField.QUARTER; + } + + if (ctx.WEEK() != null) { + return LocalDateTimeField.WEEK; + } + + if (ctx.DAY() != null) { + return LocalDateTimeField.DAY; + } + + if (ctx.HOUR() != null) { + return LocalDateTimeField.HOUR; + } + + if (ctx.MINUTE() != null) { + return LocalDateTimeField.MINUTE; + } + + if (ctx.SECOND() != null) { + return LocalDateTimeField.SECOND; + } + + throw new UnsupportedOperationException("Unsupported extract field: " + ctx.getText()); + } + + private TemporalField getTemporalField(HqlParser.ExtractFieldContext ctx) { + + if (ctx.dateOrTimeField() != null) { + + if (ctx.dateOrTimeField().DATE() != null) { + return LocalDateTimeField.DATE; + } + + if (ctx.dateOrTimeField().TIME() != null) { + return LocalDateTimeField.DATE; + } + } else if (ctx.datetimeField() != null) { + + if (ctx.datetimeField().YEAR() != null) { + return LocalDateTimeField.YEAR; + } + + if (ctx.datetimeField().MONTH() != null) { + return LocalDateTimeField.MONTH; + } + + if (ctx.datetimeField().QUARTER() != null) { + return LocalDateTimeField.QUARTER; + } + + if (ctx.datetimeField().WEEK() != null) { + return LocalDateTimeField.WEEK; + } + + if (ctx.datetimeField().DAY() != null) { + return LocalDateTimeField.DAY; + } + + if (ctx.datetimeField().HOUR() != null) { + return LocalDateTimeField.HOUR; + } + + if (ctx.datetimeField().MINUTE() != null) { + return LocalDateTimeField.MINUTE; + } + + if (ctx.datetimeField().SECOND() != null) { + return LocalDateTimeField.SECOND; + } + } else if (ctx.weekField() != null) { + + if (ctx.weekField().WEEK() != null) { + return LocalDateTimeField.WEEK; + } + + if (ctx.weekField().MONTH() != null) { + return LocalDateTimeField.MONTH; + } + + if (ctx.weekField().YEAR() != null) { + return LocalDateTimeField.YEAR; + } + } + + throw new UnsupportedOperationException("Unsupported extract field: " + ctx.getText()); + } + + @Override + @SuppressWarnings({ "rawtypes", "unchecked" }) + public Expression visitTruncFunction(HqlParser.TruncFunctionContext ctx) { + + Expression expr = visitRequired(ctx.expression().get(0)); + + if (ctx.datetimeField() != null) { + TemporalField temporalField = getTemporalField(ctx.datetimeField()); + + return cb.function("trunc", Object.class, expr, cb.literal(temporalField)); + } else if (ctx.expression().size() > 1) { + + return cb.function("trunc", Object.class, expr, visitRequired(ctx.expression().get(1))); + } + + return cb.function("trunc", Object.class, expr); + } + + @Override + public Expression visitTrimFunction(HqlParser.TrimFunctionContext ctx) { + + CriteriaBuilder.Trimspec trimSpec = null; + + HqlParser.TrimSpecificationContext tsc = ctx.trimSpecification(); + + if (tsc.LEADING() != null) { + trimSpec = CriteriaBuilder.Trimspec.LEADING; + } else if (tsc.TRAILING() != null) { + trimSpec = CriteriaBuilder.Trimspec.TRAILING; + } else if (tsc.BOTH() != null) { + trimSpec = CriteriaBuilder.Trimspec.BOTH; + } + + Expression stringLiteral = charLiteralOf(ctx.trimCharacter().STRING_LITERAL()); + Expression expression = visitRequired(ctx.expression()); + + if (trimSpec != null) { + return stringLiteral != null // + ? cb.trim(trimSpec, stringLiteral, expression) // + : cb.trim(trimSpec, expression); + } else { + return stringLiteral != null // + ? cb.trim(stringLiteral, expression) // + : cb.trim(expression); + } + } + + @Override + public Expression visitSubstringFunction(HqlParser.SubstringFunctionContext ctx) { + + Expression start = visitRequired(ctx.substringFunctionStartArgument().expression()); + + if (ctx.substringFunctionLengthArgument() != null) { + Expression length = visitRequired(ctx.substringFunctionLengthArgument().expression()); + return cb.substring(visitRequired(ctx.expression()), start, length); + } + + return cb.substring(visitRequired(ctx.expression()), start); + } + + @Override + public Expression visitLiteral(HqlParser.LiteralContext ctx) { + + if (ctx.booleanLiteral() != null) { + return visitRequired(ctx.booleanLiteral()); + } else if (ctx.JAVA_STRING_LITERAL() != null) { + return literalOf(ctx.JAVA_STRING_LITERAL()); + } else if (ctx.STRING_LITERAL() != null) { + return literalOf(ctx.STRING_LITERAL()); + } else if (ctx.numericLiteral() != null) { + return visitRequired(ctx.numericLiteral()); + } else if (ctx.temporalLiteral() != null) { + return visitRequired(ctx.temporalLiteral()); + } else if (ctx.binaryLiteral() != null) { + return visitRequired(ctx.binaryLiteral()); + } else { + return null; + } + } + + private Expression literalOf(TerminalNode node) { + + String text = node.getText(); + return cb.literal(unquoteStringLiteral(text)); + } + + private Expression charLiteralOf(TerminalNode node) { + + String text = node.getText(); + return cb.literal(text.charAt(0)); + } + + @Override + public Expression visitBooleanLiteral(HqlParser.BooleanLiteralContext ctx) { + if (ctx.TRUE() != null) { + return cb.literal(true); + } else { + return cb.literal(false); + } + } + + @Override + public Expression visitNumericLiteral(HqlParser.NumericLiteralContext ctx) { + return cb.literal(getLiteralValue(ctx)); + } + + private Number getLiteralValue(HqlParser.NumericLiteralContext ctx) { + + if (ctx.INTEGER_LITERAL() != null) { + return Integer.valueOf(getDecimals(ctx.INTEGER_LITERAL())); + } else if (ctx.LONG_LITERAL() != null) { + return Long.valueOf(getDecimals(ctx.LONG_LITERAL())); + } else if (ctx.FLOAT_LITERAL() != null) { + return Float.valueOf(getDecimals(ctx.FLOAT_LITERAL())); + } else if (ctx.DOUBLE_LITERAL() != null) { + return Double.valueOf(getDecimals(ctx.DOUBLE_LITERAL())); + } else if (ctx.BIG_INTEGER_LITERAL() != null) { + return new BigInteger(getDecimals(ctx.BIG_INTEGER_LITERAL())); + } else if (ctx.BIG_DECIMAL_LITERAL() != null) { + return new BigDecimal(getDecimals(ctx.BIG_DECIMAL_LITERAL())); + } else if (ctx.HEX_LITERAL() != null) { + return HexFormat.fromHexDigits(ctx.HEX_LITERAL().toString().substring(2)); + } + + throw new UnsupportedOperationException("Unsupported literal: " + ctx.getText()); + } + + static String getDecimals(TerminalNode input) { + + String text = input.getText(); + StringBuilder result = new StringBuilder(text.length()); + + for (int i = 0; i < text.length(); i++) { + char c = text.charAt(i); + if (Character.isDigit(c) || c == '-' || c == '+' || c == '.') { + result.append(c); + } + } + + return result.toString(); + } + + @Override + public Expression visitDateTimeLiteral(HqlParser.DateTimeLiteralContext ctx) { + + if (ctx.offsetDateTimeLiteral() != null) { + return visit(ctx.offsetDateTimeLiteral()); + } else if (ctx.localDateTimeLiteral() != null) { + return visit(ctx.localDateTimeLiteral()); + } else if (ctx.zonedDateTimeLiteral() != null) { + return visit(ctx.zonedDateTimeLiteral()); + } + + return null; + } + + @Override + public Expression visitGroupedExpression(HqlParser.GroupedExpressionContext ctx) { + return visit(ctx.expression()); + } + + @Override + public Expression visitTupleExpression(HqlParser.TupleExpressionContext ctx) { + return (Expression) cb + .tuple(ctx.expressionOrPredicate().stream().map(this::visitRequired).toArray(Expression[]::new)); + } + + @Override + public Expression visitSubqueryExpression(HqlParser.SubqueryExpressionContext ctx) { + throw new UnsupportedOperationException(String.format(UNSUPPORTED_TEMPLATE, "a subquery argument")); + } + + @Override + public Expression visitMultiplicationExpression(HqlParser.MultiplicationExpressionContext ctx) { + + Expression left = visitRequired(ctx.expression(0)); + Expression right = visitRequired(ctx.expression(1)); + + if (ctx.op.getText().equals("*")) { + return cb.prod(left, right); + } else { + return cb.quot(left, right); + } + } + + @Override + public Expression visitAdditionExpression(HqlParser.AdditionExpressionContext ctx) { + + Expression left = visitRequired(ctx.expression(0)); + Expression right = visitRequired(ctx.expression(1)); + + if (ctx.op.getText().equals("+")) { + return cb.sum(left, right); + } else { + return cb.diff(left, right); + } + } + + @Override + public Expression visitHqlConcatenationExpression(HqlParser.HqlConcatenationExpressionContext ctx) { + + Expression left = visitRequired(ctx.expression(0)); + Expression right = visitRequired(ctx.expression(1)); + + return cb.concat(left, right); + } + + @Override + public Expression visitSimplePath(HqlParser.SimplePathContext ctx) { + return QueryUtils.toExpressionRecursively((From) from, PropertyPath.from(ctx.getText(), from.getJavaType())); + } + + String getString(HqlParser.IdentifierContext context) { + + HqlParser.NakedIdentifierContext ni = context.nakedIdentifier(); + + String text = context.getText(); + if (ni != null) { + if (ni.QUOTED_IDENTIFIER() != null) { + text = unquoteIdentifier(ni.getText()); + } + } + return text; + } + + @Override + public Expression visitCaseList(HqlParser.CaseListContext ctx) { + if (ctx.simpleCaseExpression() != null) { + return visit(ctx.simpleCaseExpression()); + } else { + return visit(ctx.searchedCaseExpression()); + } + } + + @Override + public Expression visitSimpleCaseExpression(HqlParser.SimpleCaseExpressionContext ctx) { + CriteriaBuilder.SimpleCase simpleCase = cb.selectCase(visit(ctx.expressionOrPredicate(0))); + ctx.caseWhenExpressionClause().forEach(caseWhenExpressionClauseContext -> { + simpleCase.when( // + visitRequired(caseWhenExpressionClauseContext.expression()), // + visitRequired(caseWhenExpressionClauseContext.expressionOrPredicate())); + }); + if (ctx.expressionOrPredicate().size() == 2) { + simpleCase.otherwise(visitRequired(ctx.expressionOrPredicate(1))); + } + return simpleCase; + } + + @Override + public Expression visitSearchedCaseExpression(HqlParser.SearchedCaseExpressionContext ctx) { + CriteriaBuilder.Case searchedCase = cb.selectCase(); + ctx.caseWhenPredicateClause().forEach(caseWhenPredicateClauseContext -> { + searchedCase.when( // + visitRequired(caseWhenPredicateClauseContext.predicate()), // + visit(caseWhenPredicateClauseContext.expressionOrPredicate())); + }); + if (ctx.expressionOrPredicate() != null) { + searchedCase.otherwise(visit(ctx.expressionOrPredicate())); + } + return searchedCase; + } + + @Override + public Expression visitParameter(HqlParser.ParameterContext ctx) { + throw new UnsupportedOperationException(String.format(UNSUPPORTED_TEMPLATE, "a parameter argument")); + } + + @SuppressWarnings("unchecked") + private Expression visitRequired(ParseTree ctx) { + + Expression expression = visit(ctx); + + if (expression == null) { + throw new UnsupportedOperationException("No result for expression: " + ctx.getText()); + } + + return (Expression) expression; + } + + private static String unquoteIdentifier(String text) { + + int end = text.length() - 1; + assert text.charAt(0) == '`' && text.charAt(end) == '`'; + // Unquote a parsed quoted identifier and handle escape sequences + final StringBuilder sb = new StringBuilder(text.length() - 2); + for (int i = 1; i < end; i++) { + char c = text.charAt(i); + switch (c) { + case '\\': + if (i + 1 < end) { + char nextChar = text.charAt(++i); + switch (nextChar) { + case 'b': + c = '\b'; + break; + case 't': + c = '\t'; + break; + case 'n': + c = '\n'; + break; + case 'f': + c = '\f'; + break; + case 'r': + c = '\r'; + break; + case '\\': + c = '\\'; + break; + case '\'': + c = '\''; + break; + case '"': + c = '"'; + break; + case '`': + c = '`'; + break; + case 'u': + c = (char) Integer.parseInt(text.substring(i + 1, i + 5), 16); + i += 4; + break; + default: + sb.append('\\'); + c = nextChar; + break; + } + } + break; + default: + break; + } + sb.append(c); + } + return sb.toString(); + } + + private static String unquoteStringLiteral(String text) { + + int end = text.length() - 1; + char delimiter = text.charAt(0); + Assert.isTrue(delimiter == text.charAt(end), "Quoted identifier does not end with the same delimiter"); + + // Unescape the parsed literal + final StringBuilder sb = new StringBuilder(text.length() - 2); + for (int i = 1; i < end; i++) { + char c = text.charAt(i); + switch (c) { + case '\'': + if (delimiter == '\'') { + i++; + } + break; + case '"': + if (delimiter == '"') { + i++; + } + break; + default: + break; + } + sb.append(c); + } + return sb.toString(); + } + +} diff --git a/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitorUnitTests.java b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitorUnitTests.java new file mode 100644 index 000000000..cff6bea21 --- /dev/null +++ b/spring-data-jpa/src/test/java/org/springframework/data/jpa/repository/query/HqlOrderExpressionVisitorUnitTests.java @@ -0,0 +1,241 @@ +/* + * Copyright 2025 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.jpa.repository.query; + +import static org.assertj.core.api.Assertions.*; + +import jakarta.persistence.EntityManager; +import jakarta.persistence.PersistenceContext; +import jakarta.persistence.criteria.CriteriaQuery; +import jakarta.persistence.criteria.Expression; +import jakarta.persistence.criteria.Nulls; +import jakarta.persistence.criteria.Path; +import jakarta.persistence.criteria.Selection; + +import java.util.Locale; + +import org.hibernate.query.sqm.tree.select.SqmSelectStatement; +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; + +import org.springframework.data.jpa.domain.JpaSort; +import org.springframework.data.jpa.domain.sample.User; +import org.springframework.test.context.ContextConfiguration; +import org.springframework.test.context.junit.jupiter.SpringExtension; +import org.springframework.transaction.annotation.Transactional; + +/** + * Verify that {@link JpaSort#unsafe(String...)} works properly with Hibernate via {@link HqlOrderExpressionVisitor}. + * + * @author Greg Turnquist + * @author Mark Paluch + */ +@ExtendWith(SpringExtension.class) +@ContextConfiguration("classpath:application-context.xml") +@Transactional +class HqlOrderExpressionVisitorUnitTests { + + @PersistenceContext EntityManager em; + + @Test + void genericFunctions() { + + assertThat(renderOrderBy(JpaSort.unsafe("LENGTH(firstname)"), "u")) + .startsWithIgnoringCase("order by character_length(u.firstname) asc"); + assertThat(renderOrderBy(JpaSort.unsafe("char_length(firstname)"), "u")) + .startsWithIgnoringCase("order by char_length(u.firstname) asc"); + + assertThat(renderOrderBy(JpaSort.unsafe("nlssort(firstname, 'NLS_SORT = XGERMAN_DIN_AI')"), "u")) + .startsWithIgnoringCase("order by nlssort(u.firstname, 'NLS_SORT = XGERMAN_DIN_AI')"); + } + + @Test // GH-3172 + void cast() { + + assertThatExceptionOfType(UnsupportedOperationException.class) + .isThrownBy(() -> renderOrderBy(JpaSort.unsafe("cast(emailAddress as date)"), "u")); + } + + @Test // GH-3172 + void extract() { + + assertThat(renderOrderBy(JpaSort.unsafe("EXTRACT(DAY FROM createdAt)"), "u")) + .startsWithIgnoringCase("order by extract(day from u.createdAt)"); + + assertThat(renderOrderBy(JpaSort.unsafe("WEEK(createdAt)"), "u")) + .startsWithIgnoringCase("order by extract(week from u.createdAt)"); + } + + @Test // GH-3172 + void trunc() { + assertThat(renderOrderBy(JpaSort.unsafe("TRUNC(age)"), "u")).startsWithIgnoringCase("order by trunc(u.age)"); + } + + @Test // GH-3172 + void upperLower() { + assertThat(renderOrderBy(JpaSort.unsafe("upper(firstname)"), "u")) + .startsWithIgnoringCase("order by upper(u.firstname)"); + assertThat(renderOrderBy(JpaSort.unsafe("lower(firstname)"), "u")) + .startsWithIgnoringCase("order by lower(u.firstname)"); + } + + @Test // GH-3172 + void substring() { + assertThat(renderOrderBy(JpaSort.unsafe("substring(emailAddress, 0, 3)"), "u")) + .startsWithIgnoringCase("order by substring(u.emailAddress, 0, 3) asc"); + assertThat(renderOrderBy(JpaSort.unsafe("substring(emailAddress, 0)"), "u")) + .startsWithIgnoringCase("order by substring(u.emailAddress, 0) asc"); + } + + @Test // GH-3172 + void repeat() { + assertThat(renderOrderBy(JpaSort.unsafe("repeat('a', 5)"), "u")) + .startsWithIgnoringCase("order by repeat('a', 5) asc"); + } + + @Test // GH-3172 + void literals() { + + assertThat(renderOrderBy(JpaSort.unsafe("age + 1"), "u")).startsWithIgnoringCase("order by u.age + 1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1l"), "u")).startsWithIgnoringCase("order by u.age + 1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1L"), "u")).startsWithIgnoringCase("order by u.age + 1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1.1"), "u")).startsWithIgnoringCase("order by u.age + 1.1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1.1f"), "u")).startsWithIgnoringCase("order by u.age + 1.1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1.1bi"), "u")).startsWithIgnoringCase("order by u.age + 1.1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1.1bd"), "u")).startsWithIgnoringCase("order by u.age + 1.1"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 0x12"), "u")).startsWithIgnoringCase("order by u.age + 18"); + } + + @Test // GH-3172 + void arithmetic() { + + // Hibernate representation bugs, should be sum(u.age) + assertThat(renderOrderBy(JpaSort.unsafe("sum(age)"), "u")).startsWithIgnoringCase("order by sum()"); + assertThat(renderOrderBy(JpaSort.unsafe("min(age)"), "u")).startsWithIgnoringCase("order by min()"); + assertThat(renderOrderBy(JpaSort.unsafe("max(age)"), "u")).startsWithIgnoringCase("order by max()"); + + assertThat(renderOrderBy(JpaSort.unsafe("age"), "u")).startsWithIgnoringCase("order by u.age"); + assertThat(renderOrderBy(JpaSort.unsafe("age + 1"), "u")).startsWithIgnoringCase("order by u.age + 1"); + assertThat(renderOrderBy(JpaSort.unsafe("ABS(age) + 1"), "u")).startsWithIgnoringCase("order by abs(u.age) + 1"); + + assertThat(renderOrderBy(JpaSort.unsafe("neg(active)"), "u")).startsWithIgnoringCase("order by neg(u.active)"); + assertThat(renderOrderBy(JpaSort.unsafe("abs(age)"), "u")).startsWithIgnoringCase("order by abs(u.age)"); + assertThat(renderOrderBy(JpaSort.unsafe("ceiling(age)"), "u")).startsWithIgnoringCase("order by ceiling(u.age)"); + assertThat(renderOrderBy(JpaSort.unsafe("floor(age)"), "u")).startsWithIgnoringCase("order by floor(u.age)"); + assertThat(renderOrderBy(JpaSort.unsafe("round(age)"), "u")).startsWithIgnoringCase("order by round(u.age)"); + + assertThat(renderOrderBy(JpaSort.unsafe("prod(age, 1)"), "u")).startsWithIgnoringCase("order by prod(u.age, 1)"); + assertThat(renderOrderBy(JpaSort.unsafe("prod(age, age)"), "u")) + .startsWithIgnoringCase("order by prod(u.age, u.age)"); + + assertThat(renderOrderBy(JpaSort.unsafe("diff(age, 1)"), "u")).startsWithIgnoringCase("order by diff(u.age, 1)"); + assertThat(renderOrderBy(JpaSort.unsafe("quot(age, 1)"), "u")).startsWithIgnoringCase("order by quot(u.age, 1)"); + assertThat(renderOrderBy(JpaSort.unsafe("mod(age, 1)"), "u")).startsWithIgnoringCase("order by mod(u.age, 1)"); + assertThat(renderOrderBy(JpaSort.unsafe("sqrt(age)"), "u")).startsWithIgnoringCase("order by sqrt(u.age)"); + assertThat(renderOrderBy(JpaSort.unsafe("exp(age)"), "u")).startsWithIgnoringCase("order by exp(u.age)"); + assertThat(renderOrderBy(JpaSort.unsafe("ln(age)"), "u")).startsWithIgnoringCase("order by ln(u.age)"); + } + + @Test // GH-3172 + @Disabled("HHH-19075") + void trim() { + assertThat(renderOrderBy(JpaSort.unsafe("trim(leading '.' from lastname)"), "u")) + .startsWithIgnoringCase("order by repeat('a', 5) asc"); + } + + @Test // GH-3172 + void groupedExpression() { + assertThat(renderOrderBy(JpaSort.unsafe("(lastname)"), "u")).startsWithIgnoringCase("order by u.lastname"); + } + + @Test // GH-3172 + void tupleExpression() { + assertThat(renderOrderBy(JpaSort.unsafe("(firstname, lastname)"), "u")) + .startsWithIgnoringCase("order by u.firstname, u.lastname"); + } + + @Test // GH-3172 + void concat() { + assertThat(renderOrderBy(JpaSort.unsafe("firstname || lastname"), "u")) + .startsWithIgnoringCase("order by concat(u.firstname, u.lastname)"); + } + + @Test // GH-3172 + void pathBased() { + + String query = renderQuery(JpaSort.unsafe("manager.firstname"), "u"); + + assertThat(query).contains("from org.springframework.data.jpa.domain.sample.User u left join u.manager"); + assertThat(query).contains(".firstname asc nulls last"); + } + + @Test // GH-3172 + void caseSwitch() { + + assertThat(renderOrderBy(JpaSort.unsafe("case firstname when 'Oliver' then 'A' else firstname end"), "u")) + .startsWithIgnoringCase("order by case u.firstname when 'Oliver' then 'A' else u.firstname end"); + + assertThat(renderOrderBy( + JpaSort.unsafe("case firstname when 'Oliver' then 'A' when 'Joachim' then 'z' else firstname end"), "u")) + .startsWithIgnoringCase( + "order by case u.firstname when 'Oliver' then 'A' when 'Joachim' then 'z' else u.firstname end"); + + assertThat(renderOrderBy(JpaSort.unsafe("case when age < 31 then 'A' else firstname end"), "u")) + .startsWithIgnoringCase("order by case when u.age < 31 then 'A' else u.firstname end"); + + assertThat( + renderOrderBy(JpaSort.unsafe("case when firstname not in ('Oliver', 'Dave') then 'A' else firstname end"), "u")) + .startsWithIgnoringCase( + "order by case when u.firstname not in ('Oliver', 'Dave') then 'A' else u.firstname end"); + } + + private String renderOrderBy(JpaSort sort, String alias) { + + String query = renderQuery(sort, alias); + + String lowerCase = query.toLowerCase(Locale.ROOT); + int index = lowerCase.indexOf("order by"); + + if (index != -1) { + return query.substring(index); + } + + return ""; + } + + CriteriaQuery createQuery(JpaSort sort, String alias) { + + CriteriaQuery query = em.getCriteriaBuilder().createQuery(User.class); + Selection from = query.from(User.class).alias(alias); + HqlOrderExpressionVisitor extractor = new HqlOrderExpressionVisitor(em.getCriteriaBuilder(), (Path) from); + + Expression expression = extractor.createCriteriaExpression(sort.stream().findFirst().get()); + return query.select(from).orderBy(em.getCriteriaBuilder().asc(expression, Nulls.NONE)); + } + + @SuppressWarnings("rawtypes") + String renderQuery(JpaSort sort, String alias) { + + CriteriaQuery q = createQuery(sort, alias); + SqmSelectStatement s = (SqmSelectStatement) q; + + StringBuilder builder = new StringBuilder(); + s.appendHqlString(builder); + + return builder.toString(); + } +}