Add visitor to build an order expression from a JPQL order specification.

We now parse JpaSort.unsafe(…) expressions using our Query Parser and translate the parsed tree into a CriteriaQuery Expression except for CAST, TREAT and subqueries.

Closes #3172
Original pull request: #3187
This commit is contained in:
Greg L. Turnquist
2025-01-27 15:54:33 +01:00
committed by Mark Paluch
parent 1096088c41
commit b61c51d8c4
2 changed files with 981 additions and 0 deletions

View File

@@ -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<Expression<?>> {
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<Comparable> left = visitRequired(ctx.expression(0));
Expression<Comparable> 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<Comparable> condition = visitRequired(ctx.expression(0));
Expression<Comparable> lower = visitRequired(ctx.expression(1));
Expression<Comparable> 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<? extends Collection<?>>) condition);
} else {
return cb.isNotEmpty((Expression<? extends Collection<?>>) condition);
}
}
if (ctx.TRUE() != null) {
if (ctx.NOT() == null) {
return cb.isTrue((Expression<Boolean>) condition);
} else {
return cb.isFalse((Expression<Boolean>) condition);
}
}
if (ctx.FALSE() != null) {
if (ctx.NOT() == null) {
return cb.isFalse((Expression<Boolean>) condition);
} else {
return cb.isTrue((Expression<Boolean>) condition);
}
}
return null;
}
@Override
public Expression<?> visitStringPatternMatching(HqlParser.StringPatternMatchingContext ctx) {
Expression<String> condition = visitRequired(ctx.expression(0));
Expression<String> match = visitRequired(ctx.expression(1));
Expression<Character> 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<Object> 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<?, ? extends Temporal> 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<?, ? extends Temporal> 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<Character> stringLiteral = charLiteralOf(ctx.trimCharacter().STRING_LITERAL());
Expression<String> 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<Integer> start = visitRequired(ctx.substringFunctionStartArgument().expression());
if (ctx.substringFunctionLengthArgument() != null) {
Expression<Integer> 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<String> literalOf(TerminalNode node) {
String text = node.getText();
return cb.literal(unquoteStringLiteral(text));
}
private Expression<Character> 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<Number> left = visitRequired(ctx.expression(0));
Expression<Number> 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<Number> left = visitRequired(ctx.expression(0));
Expression<Number> 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<String> left = visitRequired(ctx.expression(0));
Expression<String> 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<Object, Object> 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<Object> 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 <T> Expression<T> visitRequired(ParseTree ctx) {
Expression<?> expression = visit(ctx);
if (expression == null) {
throw new UnsupportedOperationException("No result for expression: " + ctx.getText());
}
return (Expression<T>) 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();
}
}

View File

@@ -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<User> createQuery(JpaSort sort, String alias) {
CriteriaQuery<User> query = em.getCriteriaBuilder().createQuery(User.class);
Selection<User> 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<User> q = createQuery(sort, alias);
SqmSelectStatement s = (SqmSelectStatement) q;
StringBuilder builder = new StringBuilder();
s.appendHqlString(builder);
return builder.toString();
}
}