PredicateSpecificationEvaluator.java
package com.reallifedeveloper.tools.test.database.inmemory;
import java.lang.reflect.Field;
import java.lang.reflect.InvocationHandler;
import java.lang.reflect.Method;
import java.lang.reflect.Proxy;
import java.math.BigDecimal;
import java.util.AbstractMap;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collection;
import java.util.HashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Objects;
import java.util.Optional;
import java.util.concurrent.ConcurrentHashMap;
import java.util.function.Function;
import org.checkerframework.checker.nullness.qual.Nullable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.data.jpa.domain.PredicateSpecification;
import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
import jakarta.persistence.criteria.CriteriaBuilder;
import jakarta.persistence.criteria.Expression;
import jakarta.persistence.criteria.From;
import jakarta.persistence.criteria.Join;
import jakarta.persistence.criteria.JoinType;
import jakarta.persistence.criteria.MapJoin;
import jakarta.persistence.criteria.Path;
import jakarta.persistence.criteria.Predicate;
import lombok.Getter;
import lombok.experimental.Accessors;
/**
* An evaluator of {@code org.springframework.data.jpa.domain.PredicateSpecification} instances.
* <p>
* This implementation is maintained in a single, large, file. This is because, a) the functionatlity has a clear focus with a small public
* API, and b) the code was written for the main part by ChatGPT, and I want to make it clear what code was generated by AI.
*
* @param <T> the type of entity for which the specification
*
* @see <a href= "https://docs.spring.io/spring-data/jpa/reference/jpa/specifications.html#predicate-specification">Spring Data JPA
* documentation</a>
*
* @author ChatGPT, RealLifeDeveloper
*/
@SuppressWarnings({ "InnerTypeLast", "PMD" })
@SuppressFBWarnings(value = "CRLF_INJECTION_LOGS", justification = "This code is only intended for testing")
public final class PredicateSpecificationEvaluator<T> {
private static final Logger LOG = LoggerFactory.getLogger(PredicateSpecificationEvaluator.class);
private final Map<String, RegisteredFunction> functions = new HashMap<>();
/**
* Checks if the given entity matches the given specification.
*
* @param specification the {@code PredicateSpecification} to use
* @param entity the entity to check
*
* @return {@code true} if {@code specification} matches {@code entity}, {@code false} otherwise
*/
public boolean matches(PredicateSpecification<T> specification, T entity) {
Objects.requireNonNull(specification);
Objects.requireNonNull(entity);
RecordingState state = new RecordingState();
From<?, T> root = Proxies.root(state);
CriteriaBuilder cb = Proxies.criteriaBuilder();
Predicate predicate = specification.toPredicate(root, cb);
// PredicateSpecification.unrestricted() and similar specifications are represented by a null predicate.
if (predicate == null) {
return true;
}
Expr expr = Proxies.expressionOf(predicate);
LOG.debug("matches: expr={}, entity={}", expr, entity);
LOG.trace("matches: joins={}", state.joins());
List<EvaluationContext> rows = expandJoins(state, entity);
LOG.trace("matches: rows={}", rows);
return rows.stream().anyMatch(context -> evaluateBoolean(expr, context) == Truth.TRUE);
}
/**
* Filters a collection of entities based on if they match the given specification or not.
*
* @param specification the {@code PredicateSpecification} to use
* @param entities the collection of entities to filter
*
* @return a list of the entities from the {@code entities} collection that match {@code specification}
*/
public List<T> filter(PredicateSpecification<T> specification, Collection<T> entities) {
return entities.stream().filter(entity -> matches(specification, entity)).toList();
}
/**
* Registers a function that can be used by {@code CriteriaBuilder.function}.
* <p>
* An example of registering a function:
*
* <pre>{@code
* evaluator.registerFunction("concat_with_separator", String.class, arguments -> {
* String separator = (String) arguments.get(0);
* String left = (String) arguments.get(1);
* String right = (String) arguments.get(2);
* if (separator == null || left == null || right == null) {
* return null;
* }
* return left + separator + right;
* });
* }</pre>
*
* @param name the name of the function to register
* @param resultType the class representing the return type of the function to register
* @param parameterTypes a list of class objects representing the type of the parameters
* @param function the function to register
* @param <R> the return type of the function to register
*
* @return the {@code PredicationSpecificationEvaluator} itself, in order to support nested (fluent) calls
*/
public <R> PredicateSpecificationEvaluator<T> registerFunction(String name, Class<R> resultType, List<Class<?>> parameterTypes,
EvaluationFunction function) {
Objects.requireNonNull(name);
Objects.requireNonNull(resultType);
Objects.requireNonNull(parameterTypes);
Objects.requireNonNull(function);
functions.put(normalizeFunctionName(name), new RegisteredFunction(resultType, List.copyOf(parameterTypes), function));
return this;
}
/**
* A function that can be registered using the {@link #registerFunction(String, Class, List, EvaluationFunction)} method.
*/
@FunctionalInterface
public interface EvaluationFunction {
/**
* Applies the function to its arguments.
*
* @param arguments the list of arguments
*
* @return the return values of the function
*/
Object apply(List<Object> arguments);
}
/**
* Registers a single-parameter function thet can be used by {@code CriteriaBuilder.function}.
* <p>
* This is a convenience method so that you can register functions with a single parameter like this:
*
* <pre>{@code
* evaluator.registerFunction("lower", String.class, String.class, value -> value == null ? null : value.toLowerCase(Locale.ROOT));
* }</pre>
*
* @param name the name of the function to register
* @param resultType the class representing the return type of the function to register
* @param argumentType the class representing the type of the only function parameter
* @param function the function to register
* @param <A> the type of the only function parameter
* @param <R> the return type of the function to register
*
* @return the {@code PredicationSpecificationEvaluator} itself, in order to support nested (fluent) calls
*/
public <A, R> PredicateSpecificationEvaluator<T> registerFunction(String name, Class<R> resultType, Class<A> argumentType,
Function<A, R> function) {
return registerFunction(name, resultType, List.of(argumentType), arguments -> function.apply(argumentType.cast(arguments.get(0))));
}
// ============================================================
// Expression Model
// ============================================================
private sealed interface Expr {
}
private record Constant(@Nullable Object value) implements Expr {
}
private sealed interface PathExpression extends Expr {
}
private record PathExpr(Source source, List<String> attributes) implements PathExpression {
}
private record AttributePath(Expr parent, String attribute) implements PathExpression {
}
private record MapKeyExpr(Source source) implements PathExpression {
}
private record MapValueExpr(Source source) implements PathExpression {
}
private record Equal(Expr left, Expr right) implements Expr {
}
private record NotEqual(Expr left, Expr right) implements Expr {
}
private record GreaterThan(Expr left, Expr right) implements Expr {
}
private record GreaterThanOrEqual(Expr left, Expr right) implements Expr {
}
private record LessThan(Expr left, Expr right) implements Expr {
}
private record LessThanOrEqual(Expr left, Expr right) implements Expr {
}
private record IsNull(Expr expression) implements Expr {
}
private record IsNotNull(Expr expression) implements Expr {
}
private record And(List<Expr> expressions) implements Expr {
}
private record Or(List<Expr> expressions) implements Expr {
}
private record Not(Expr expression) implements Expr {
}
private record Product(Expr left, Expr right) implements Expr {
}
private record In(Expr expression, List<Expr> values) implements Expr {
}
private record TrueExpr() implements Expr {
}
private record FalseExpr() implements Expr {
}
private record MapEntryExpr(Source source) implements Expr {
}
private record IsEmpty(Expr expression) implements Expr {
}
private record IsNotEmpty(Expr expression) implements Expr {
}
private record FunctionCall(String name, Class<?> resultType, List<Expr> arguments) implements Expr {
}
private interface ExpressionProvider {
Expr expression();
}
// ============================================================
// Sources and Joins
// ============================================================
private sealed interface Source {
}
private record RootSource() implements Source {
}
private record JoinSource(int id) implements Source {
}
private enum JoinKind {
STANDARD, MAP
}
private record JoinDefinition(int id, Source parent, String attribute, JoinType joinType, JoinKind kind) {
}
@Getter
@Accessors(fluent = true)
private static final class RecordingState {
private int nextJoinId;
private final List<JoinDefinition> joins = new ArrayList<>();
JoinSource addJoin(Source parent, String attribute, JoinType joinType) {
return addJoin(parent, attribute, joinType, JoinKind.STANDARD);
}
JoinSource addMapJoin(Source parent, String attribute, JoinType joinType) {
return addJoin(parent, attribute, joinType, JoinKind.MAP);
}
private JoinSource addJoin(Source parent, String attribute, JoinType joinType, JoinKind kind) {
int id = nextJoinId++;
joins.add(new JoinDefinition(id, parent, attribute, joinType, kind));
return new JoinSource(id);
}
}
// ============================================================
// Evaluation Context
// ============================================================
private record EvaluationContext(Object root, Map<Integer, Object> joins) {
EvaluationContext bind(int joinId, @Nullable Object value) {
Map<Integer, Object> copy = new HashMap<>(joins);
copy.put(joinId, value);
return new EvaluationContext(root, copy);
}
@Nullable
Object source(Source source) {
return switch (source) {
case RootSource ignored -> root;
case JoinSource join -> joins.get(join.id());
};
}
}
// ============================================================
// Join Expansion
// ============================================================
private List<EvaluationContext> expandJoins(RecordingState state, Object entity) {
List<EvaluationContext> rows = List.of(new EvaluationContext(entity, Map.of()));
for (JoinDefinition join : state.joins) {
List<EvaluationContext> next = new ArrayList<>();
for (EvaluationContext row : rows) {
Object parent = resolveJoinParent(row, join.parent());
Object value = parent == null ? null : PropertyAccess.read(parent, join.attribute());
List<?> joinedValues = normalizeJoinValue(value, join.kind());
if (joinedValues.isEmpty()) {
if (join.joinType() == JoinType.LEFT) {
next.add(row.bind(join.id(), null));
}
continue;
}
for (Object joinedValue : joinedValues) {
next.add(row.bind(join.id(), joinedValue));
}
}
rows = next;
}
return rows;
}
private record MapBinding(Object key, Object value) {
}
private static @Nullable Object resolveJoinParent(EvaluationContext context, Source source) {
Object value = context.source(source);
if (value instanceof MapBinding mapBinding) {
return mapBinding.value();
}
return value;
}
private static List<?> normalizeJoinValue(@Nullable Object value, JoinKind kind) {
if (value == null) {
return List.of();
}
if (kind == JoinKind.MAP) {
if (!(value instanceof Map<?, ?> map)) {
throw new IllegalArgumentException("Map join requires a Map value, but got " + value.getClass().getName());
}
return map.entrySet().stream().map(entry -> new MapBinding(entry.getKey(), entry.getValue())).toList();
}
if (value instanceof Collection<?> collection) {
return new ArrayList<>(collection);
}
return List.of(value);
}
// ============================================================
// Expression Evaluation
// ============================================================
private enum Truth {
TRUE, FALSE, UNKNOWN
}
private Truth evaluateBoolean(Expr expr, EvaluationContext context) {
LOG.trace("evaluateBoolean: expr={}", expr);
return switch (expr) {
case TrueExpr ignored -> Truth.TRUE;
case FalseExpr ignored -> Truth.FALSE;
case Equal e -> equal(evaluateValue(e.left(), context), evaluateValue(e.right(), context));
case NotEqual e -> not(equal(evaluateValue(e.left(), context), evaluateValue(e.right(), context)));
case GreaterThan e -> compare(e.left(), e.right(), context, result -> result > 0);
case GreaterThanOrEqual e -> compare(e.left(), e.right(), context, result -> result >= 0);
case LessThan e -> compare(e.left(), e.right(), context, result -> result < 0);
case LessThanOrEqual e -> compare(e.left(), e.right(), context, result -> result <= 0);
case IsNull e -> evaluateValue(e.expression(), context) == null ? Truth.TRUE : Truth.FALSE;
case IsNotNull e -> evaluateValue(e.expression(), context) != null ? Truth.TRUE : Truth.FALSE;
case And e -> evaluateAnd(e.expressions(), context);
case Or e -> evaluateOr(e.expressions(), context);
case Not e -> not(evaluateBoolean(e.expression(), context));
case In e -> evaluateIn(e, context);
case IsEmpty e -> evaluateIsEmpty(e, context);
case IsNotEmpty e -> evaluateIsNotEmpty(e, context);
default -> throw new IllegalArgumentException("Not a boolean expression: " + expr);
};
}
private @Nullable Object evaluateValue(Expr expr, EvaluationContext context) {
LOG.trace("evaluateValue: expr={}", expr);
return switch (expr) {
case Constant constant -> constant.value();
case PathExpr path -> evaluatePath(path, context);
case AttributePath path -> {
Object parent = evaluateValue(path.parent(), context);
yield parent == null ? null : PropertyAccess.read(parent, path.attribute());
}
case MapKeyExpr mapKey -> mapBinding(mapKey.source(), context).key();
case MapValueExpr mapValue -> mapBinding(mapValue.source(), context).value();
case MapEntryExpr mapEntry -> {
MapBinding binding = mapBinding(mapEntry.source(), context);
yield new AbstractMap.SimpleImmutableEntry<>(binding.key(), binding.value());
}
case Product product -> multiply(evaluateValue(product.left(), context), evaluateValue(product.right(), context));
case FunctionCall function -> evaluateFunction(function, context);
default -> throw new IllegalArgumentException("Not a value expression: " + expr);
};
}
@SuppressWarnings("noReturnNull")
private @Nullable Object evaluatePath(PathExpr path, EvaluationContext context) {
Object current = context.source(path.source());
if (current instanceof MapBinding mapBinding) {
current = mapBinding.value();
}
for (String attribute : path.attributes()) {
if (current == null) {
return null;
}
current = PropertyAccess.read(current, attribute);
}
return current;
}
private static MapBinding mapBinding(Source source, EvaluationContext context) {
Object value = context.source(source);
if (!(value instanceof MapBinding binding)) {
throw new IllegalStateException("Expected map binding for " + source);
}
return binding;
}
private static Truth equal(@Nullable Object left, @Nullable Object right) {
// SQL semantics:
// NULL = anything => UNKNOWN.
if (left == null || right == null) {
return Truth.UNKNOWN;
}
return Objects.equals(left, right) ? Truth.TRUE : Truth.FALSE;
}
private Truth evaluateAnd(List<Expr> expressions, EvaluationContext context) {
LOG.trace("evaluateAnd: expressions={}", expressions);
Truth result = Truth.TRUE;
for (Expr expr : expressions) {
LOG.trace("evaluateAnd: expr={}", expr);
Truth value = evaluateBoolean(expr, context);
LOG.trace("evaluateAnd: value={}", value);
if (value == Truth.FALSE) {
return Truth.FALSE;
}
if (value == Truth.UNKNOWN) {
result = Truth.UNKNOWN;
}
}
return result;
}
private Truth evaluateOr(List<Expr> expressions, EvaluationContext context) {
LOG.trace("evaluateOr: expressions={}", expressions);
Truth result = Truth.FALSE;
for (Expr expr : expressions) {
LOG.trace("evaluateOr: expr={}", expr);
Truth value = evaluateBoolean(expr, context);
LOG.trace("evaluateOr: value={}", value);
if (value == Truth.TRUE) {
return Truth.TRUE;
}
if (value == Truth.UNKNOWN) {
result = Truth.UNKNOWN;
}
}
return result;
}
private static Truth not(Truth truth) {
return switch (truth) {
case TRUE -> Truth.FALSE;
case FALSE -> Truth.TRUE;
case UNKNOWN -> Truth.UNKNOWN;
};
}
private Truth evaluateIn(In in, EvaluationContext context) {
Object testedValue = evaluateValue(in.expression(), context);
boolean unknown = false;
for (Expr valueExpr : in.values()) {
Object candidate = evaluateValue(valueExpr, context);
Truth comparison = equal(testedValue, candidate);
if (comparison == Truth.TRUE) {
return Truth.TRUE;
}
if (comparison == Truth.UNKNOWN) {
unknown = true;
}
}
return unknown ? Truth.UNKNOWN : Truth.FALSE;
}
private Truth evaluateIsEmpty(IsEmpty expression, EvaluationContext context) {
Object value = evaluateValue(expression.expression(), context);
if (value == null) {
return Truth.UNKNOWN;
}
if (!(value instanceof Collection<?> collection)) {
throw new IllegalArgumentException("isEmpty() requires a Collection, but got " + value.getClass().getName());
}
return collection.isEmpty() ? Truth.TRUE : Truth.FALSE;
}
private Truth evaluateIsNotEmpty(IsNotEmpty expression, EvaluationContext context) {
Object value = evaluateValue(expression.expression(), context);
if (value == null) {
return Truth.UNKNOWN;
}
if (!(value instanceof Collection<?> collection)) {
throw new IllegalArgumentException("isNotEmpty() requires a Collection, but got " + value.getClass().getName());
}
return collection.isEmpty() ? Truth.FALSE : Truth.TRUE;
}
private Truth compare(Expr leftExpr, Expr rightExpr, EvaluationContext context, java.util.function.IntPredicate condition) {
Object left = evaluateValue(leftExpr, context);
Object right = evaluateValue(rightExpr, context);
LOG.trace("compare: left={}, right={}", left, right);
if (left == null || right == null) {
return Truth.UNKNOWN;
}
int result = compareValues(left, right);
return condition.test(result) ? Truth.TRUE : Truth.FALSE;
}
@SuppressWarnings({ "rawtypes", "unchecked" })
private static int compareValues(Object left, Object right) {
if (left instanceof Number l && right instanceof Number r) {
return toBigDecimal(l).compareTo(toBigDecimal(r));
}
if (left instanceof Comparable comparable) {
return comparable.compareTo(right);
}
throw new IllegalArgumentException("Cannot compare " + left.getClass().getName() + " and " + right.getClass().getName());
}
@SuppressWarnings("noReturnNull")
private static @Nullable Object multiply(@Nullable Object left, @Nullable Object right) {
LOG.trace("multiply: left={}, right={}", left, right);
if (left == null || right == null) {
return null;
}
if (!(left instanceof Number l) || !(right instanceof Number r)) {
throw new IllegalArgumentException("prod() requires numeric operands");
}
return toBigDecimal(l).multiply(toBigDecimal(r));
}
private static BigDecimal toBigDecimal(Number value) {
if (value instanceof BigDecimal bd) {
return bd;
}
return new BigDecimal(value.toString());
}
// ============================================================
// Criteria API Proxy Implementation
// ============================================================
private static final class Proxies {
private Proxies() {
}
@SuppressWarnings("unchecked")
static <T> From<?, T> root(RecordingState state) {
return (From<?, T>) Proxy.newProxyInstance(From.class.getClassLoader(), new Class<?>[] { From.class },
new FromHandler(state, new RootSource()));
}
static CriteriaBuilder criteriaBuilder() {
return (CriteriaBuilder) Proxy.newProxyInstance(CriteriaBuilder.class.getClassLoader(),
new Class<?>[] { CriteriaBuilder.class }, new CriteriaBuilderHandler());
}
static Expr expressionOf(Object value) {
LOG.trace("expressionOf: value={}", value);
if (value == null) {
return new Constant(null);
}
if (Proxy.isProxyClass(value.getClass())) {
InvocationHandler handler = Proxy.getInvocationHandler(value);
if (handler instanceof ExpressionProvider provider) {
return provider.expression();
}
}
return new Constant(value);
}
@SuppressWarnings("unchecked")
static <X> Path<X> path(PathExpression expression) {
return (Path<X>) Proxy.newProxyInstance(Path.class.getClassLoader(), new Class<?>[] { Path.class },
new ExpressionHandler(expression));
}
static Predicate predicate(Expr expr) {
return (Predicate) Proxy.newProxyInstance(Predicate.class.getClassLoader(), new Class<?>[] { Predicate.class },
new ExpressionHandler(expr));
}
@SuppressWarnings("unchecked")
static <X> Expression<X> expression(Expr expr) {
return (Expression<X>) Proxy.newProxyInstance(Expression.class.getClassLoader(), new Class<?>[] { Expression.class },
new ExpressionHandler(expr));
}
@SuppressWarnings("unchecked")
static <X, Y> Join<X, Y> join(RecordingState state, JoinSource source) {
return (Join<X, Y>) Proxy.newProxyInstance(Join.class.getClassLoader(), new Class<?>[] { Join.class },
new FromHandler(state, source));
}
@SuppressWarnings("unchecked")
static <X, K, V> MapJoin<X, K, V> mapJoin(RecordingState state, JoinSource source) {
return (MapJoin<X, K, V>) Proxy.newProxyInstance(MapJoin.class.getClassLoader(), new Class<?>[] { MapJoin.class },
new MapJoinHandler(state, source));
}
@SuppressWarnings("unchecked")
static <T> CriteriaBuilder.In<T> in(Expr expression) {
return (CriteriaBuilder.In<T>) Proxy.newProxyInstance(CriteriaBuilder.In.class.getClassLoader(),
new Class<?>[] { CriteriaBuilder.In.class }, new InHandler(expression));
}
}
// ============================================================
// From / Join Proxy
// ============================================================
private static final class FromHandler implements InvocationHandler, ExpressionProvider {
private final RecordingState state;
private final Source source;
private FromHandler(RecordingState state, Source source) {
this.state = state;
this.source = source;
}
@Override
public Expr expression() {
return new PathExpr(source, List.of());
}
@Override
public Object invoke(Object proxy, Method method, Object[] args) {
String name = method.getName();
if (isGet(method, args)) {
return navigatePath(expression(), (String) args[0]);
}
if (name.equals("join") && args != null && args.length >= 1 && args[0] instanceof String attribute) {
JoinType joinType = args.length >= 2 && args[1] instanceof JoinType jt ? jt : JoinType.INNER;
JoinSource join = state.addJoin(source, attribute, joinType);
return Proxies.join(state, join);
}
if (method.getName().equals("joinMap") && args != null && args.length >= 1 && args[0] instanceof String attribute) {
JoinType joinType = args.length >= 2 && args[1] instanceof JoinType jt ? jt : JoinType.INNER;
JoinSource join = state.addMapJoin(source, attribute, joinType);
return Proxies.mapJoin(state, join);
}
return objectMethodOrUnsupported(proxy, method, args, "From[" + source + "]");
}
}
private static boolean isGet(Method method, Object[] args) {
return method.getName().equals("get") && args != null && args.length == 1 && args[0] instanceof String;
}
private static Object navigatePath(Expr parent, String attribute) {
return Proxies.path(new AttributePath(parent, attribute));
}
// ============================================================
// MapJoin Proxy
// ============================================================
private static final class MapJoinHandler implements InvocationHandler, ExpressionProvider {
private final RecordingState state;
private final JoinSource source;
private MapJoinHandler(RecordingState state, JoinSource source) {
this.state = state;
this.source = source;
}
@Override
public Expr expression() {
return new PathExpr(source, List.of());
}
@Override
public Object invoke(Object proxy, Method method, Object[] args) {
String name = method.getName();
if (isGet(method, args)) {
return navigatePath(expression(), (String) args[0]);
}
if (name.equals("key") && method.getParameterCount() == 0) {
return Proxies.path(new MapKeyExpr(source));
}
if (name.equals("value") && method.getParameterCount() == 0) {
return Proxies.path(new MapValueExpr(source));
}
if (name.equals("entry") && method.getParameterCount() == 0) {
return Proxies.expression(new MapEntryExpr(source));
}
/*
* A MapJoin is also a From, so allow nested joins against the map value.
*/
if (name.equals("join") && args != null && args.length >= 1 && args[0] instanceof String attribute) {
JoinType joinType = args.length >= 2 && args[1] instanceof JoinType jt ? jt : JoinType.INNER;
JoinSource join = state.addJoin(source, attribute, joinType);
return Proxies.join(state, join);
}
if (name.equals("joinMap") && args != null && args.length >= 1 && args[0] instanceof String attribute) {
JoinType joinType = args.length >= 2 && args[1] instanceof JoinType jt ? jt : JoinType.INNER;
JoinSource join = state.addMapJoin(source, attribute, joinType);
return Proxies.mapJoin(state, join);
}
return objectMethodOrUnsupported(proxy, method, args, "MapJoin[" + source + "]");
}
}
// ============================================================
// Path / Expression / Predicate Proxy
// ============================================================
private static final class ExpressionHandler implements InvocationHandler, ExpressionProvider {
private final Expr expression;
private ExpressionHandler(Expr expression) {
LOG.trace("ExpressionHandler: expression={}", expression);
this.expression = expression;
}
@Override
public Expr expression() {
return expression;
}
@Override
public Object invoke(Object proxy, Method method, Object[] args) {
/*
* Path.get("attribute")
*/
if (isGet(method, args)) {
return navigatePath(expression, (String) args[0]);
}
/*
* Expression.in(...)
*
* Handles:
*
* path.in("A", "B") path.in(List.of("A", "B")) path.in(expr1, expr2)
*/
if (method.getName().equals("in")) {
return handleIn(expression, method, args);
}
/*
* Predicate.not()
*
* Note that cb.not(predicate) is handled by CriteriaBuilderHandler instead.
*/
if (method.getName().equals("not") && method.getParameterCount() == 0 && proxy instanceof Predicate) {
return Proxies.predicate(new Not(expression));
}
return objectMethodOrUnsupported(proxy, method, args, expression.toString());
}
}
private static Object handleIn(Expr expression, Method method, Object[] args) {
if (args == null || args.length != 1) {
throw new UnsupportedOperationException("Unsupported Expression.in() overload: " + method);
}
LOG.trace("handleIn: expression={}, method={}, args={}", expression, method, Arrays.asList(args));
Object argument = args[0];
List<Expr> values;
if (argument instanceof Collection<?> collection) {
values = collection.stream().map(Proxies::expressionOf).toList();
} else if (argument instanceof Object[] array) {
values = Arrays.stream(array).map(Proxies::expressionOf).toList();
} else {
values = List.of(Proxies.expressionOf(argument));
}
return Proxies.predicate(new In(expression, values));
}
// ============================================================
// CriteriaBuilder Proxy
// ============================================================
private static final class CriteriaBuilderHandler implements InvocationHandler {
@Override
public Object invoke(Object proxy, Method method, Object[] args) {
String name = method.getName();
return switch (name) {
case "equal" -> Proxies.predicate(new Equal(expr(args[0]), expr(args[1])));
case "notEqual" -> Proxies.predicate(new NotEqual(expr(args[0]), expr(args[1])));
case "greaterThan" -> Proxies.predicate(new GreaterThan(expr(args[0]), expr(args[1])));
case "greaterThanOrEqualTo" -> Proxies.predicate(new GreaterThanOrEqual(expr(args[0]), expr(args[1])));
case "lessThan" -> Proxies.predicate(new LessThan(expr(args[0]), expr(args[1])));
case "lessThanOrEqualTo" -> Proxies.predicate(new LessThanOrEqual(expr(args[0]), expr(args[1])));
case "isNull" -> Proxies.predicate(new IsNull(expr(args[0])));
case "isNotNull" -> Proxies.predicate(new IsNotNull(expr(args[0])));
case "and" -> Proxies.predicate(new And(predicateArguments(args)));
case "or" -> Proxies.predicate(new Or(predicateArguments(args)));
case "not" -> Proxies.predicate(new Not(expr(args[0])));
case "conjunction" -> Proxies.predicate(new TrueExpr());
case "disjunction" -> Proxies.predicate(new FalseExpr());
case "literal" -> Proxies.expression(new Constant(args[0]));
case "prod" -> Proxies.expression(new Product(expr(args[0]), expr(args[1])));
case "in" -> Proxies.in(expr(args[0]));
case "isEmpty" -> Proxies.predicate(new IsEmpty(expr(args[0])));
case "isNotEmpty" -> Proxies.predicate(new IsNotEmpty(expr(args[0])));
case "function" -> handleFunction(args);
default -> objectMethodOrUnsupported(proxy, method, args, "CriteriaBuilder");
};
}
private static Expr expr(Object value) {
return Proxies.expressionOf(value);
}
private static List<Expr> predicateArguments(Object[] args) {
if (args == null || args.length == 0) {
return List.of();
}
/*
* For CriteriaBuilder.and(Predicate...) reflection sees one Predicate[] argument.
*
* For and(Expression<Boolean>, Expression<Boolean>) it sees two arguments.
*/
if (args.length == 1 && args[0] instanceof Object[] array) {
return Arrays.stream(array).map(Proxies::expressionOf).toList();
}
return Arrays.stream(args).map(Proxies::expressionOf).toList();
}
}
// ============================================================
// CriteriaBuilder.In Proxy
// ============================================================
private static final class InHandler implements InvocationHandler, ExpressionProvider {
private final Expr expression;
private final List<Expr> values = new ArrayList<>();
private InHandler(Expr expression) {
this.expression = expression;
}
@Override
public Expr expression() {
return new In(expression, List.copyOf(values));
}
@Override
public Object invoke(Object proxy, Method method, Object[] args) {
if (method.getName().equals("value") && args != null && args.length == 1) {
values.add(Proxies.expressionOf(args[0]));
// CriteriaBuilder.In.value() returns itself.
return proxy;
}
if (method.getName().equals("not")) {
return Proxies.predicate(new Not(expression()));
}
return objectMethodOrUnsupported(proxy, method, args, expression().toString());
}
}
// ============================================================
// Bean / Property Access
// ============================================================
private static final class PropertyAccess {
private static final Map<Key, Accessor> CACHE = new ConcurrentHashMap<>();
@SuppressWarnings("noReturnNull")
static @Nullable Object read(@Nullable Object target, String property) {
LOG.trace("PropertyAccess.read: target={}, property={}", target, property);
if (target == null) {
return null;
}
Accessor accessor = CACHE.computeIfAbsent(new Key(target.getClass(), property), PropertyAccess::findAccessor);
Object value = accessor.read(target);
return unwrapOptionalIfNecessary(value);
}
private static @Nullable Object unwrapOptionalIfNecessary(Object value) {
if (value instanceof Optional<?> optional) {
return optional.orElse(null);
}
return value;
}
private static Accessor findAccessor(Key key) {
Class<?> type = key.type();
String property = key.property();
/*
* Records and ordinary methods whose name exactly equals the property.
*/
try {
Method method = type.getMethod(property);
if (method.getParameterCount() == 0) {
return new MethodAccessor(method);
}
} catch (NoSuchMethodException ignored) {
// Empty
}
String capitalized = Character.toUpperCase(property.charAt(0)) + property.substring(1);
for (String name : List.of("get" + capitalized, "is" + capitalized)) {
try {
Method method = type.getMethod(name);
if (method.getParameterCount() == 0) {
return new MethodAccessor(method);
}
} catch (NoSuchMethodException ignored) {
// Empty
}
}
Class<?> current = type;
while (current != null) {
try {
Field field = current.getDeclaredField(property);
field.trySetAccessible();
return new FieldAccessor(field);
} catch (NoSuchFieldException ignored) {
current = current.getSuperclass();
}
}
throw new IllegalArgumentException("No readable property '" + property + "' on " + type.getName());
}
private record Key(Class<?> type, String property) {
}
private interface Accessor {
Object read(Object target);
}
private record MethodAccessor(Method method) implements Accessor {
@Override
public Object read(Object target) {
try {
return method.invoke(target);
} catch (ReflectiveOperationException e) {
throw new IllegalStateException(e);
}
}
}
private record FieldAccessor(Field field) implements Accessor {
@Override
public Object read(Object target) {
try {
return field.get(target);
} catch (IllegalAccessException e) {
throw new IllegalStateException(e);
}
}
}
}
// ============================================================
// Function Support
// ============================================================
private record RegisteredFunction(Class<?> resultType, List<Class<?>> argumentTypes, EvaluationFunction function) {
}
private static String normalizeFunctionName(String name) {
return name.toLowerCase(Locale.ROOT);
}
@SuppressWarnings("MagicNumber")
private static Object handleFunction(Object[] args) {
if (args == null || args.length != 3 || !(args[0] instanceof String name) || !(args[1] instanceof Class<?> resultType)
|| !(args[2] instanceof Object[] functionArguments)) {
throw new UnsupportedOperationException("Unsupported CriteriaBuilder.function() invocation");
}
List<Expr> arguments = Arrays.stream(functionArguments).map(Proxies::expressionOf).toList();
return Proxies.expression(new FunctionCall(name, resultType, arguments));
}
private Object evaluateFunction(FunctionCall functionCall, EvaluationContext context) {
RegisteredFunction registered = functions.get(normalizeFunctionName(functionCall.name()));
if (registered == null) {
throw new UnsupportedOperationException("No evaluator function registered for CriteriaBuilder.function(" + functionCall.name()
+ ", " + functionCall.resultType().getName() + ")");
}
verifyCompatibleResultType(functionCall, registered);
List<Object> arguments = functionCall.arguments().stream().map(argument -> evaluateValue(argument, context)).toList();
validateFunctionArguments(functionCall.name(), registered, arguments);
Object result = registered.function().apply(arguments);
verifyFunctionResult(functionCall, result);
return result;
}
private static void verifyCompatibleResultType(FunctionCall call, RegisteredFunction registered) {
if (!wrap(call.resultType()).isAssignableFrom(wrap(registered.resultType()))) {
throw new IllegalArgumentException("Function '%s' was requested with result type %s, but is registered with result type %s"
.formatted(call.name(), call.resultType().getName(), registered.resultType().getName()));
}
}
private static void validateFunctionArguments(String functionName, RegisteredFunction function, List<Object> arguments) {
List<Class<?>> expectedTypes = function.argumentTypes();
if (arguments.size() != expectedTypes.size()) {
throw new IllegalArgumentException(
"Function '%s' expected %d arguments but received %d".formatted(functionName, expectedTypes.size(), arguments.size()));
}
for (int i = 0; i < arguments.size(); i++) {
Object argument = arguments.get(i);
Class<?> expectedType = wrap(expectedTypes.get(i));
if (argument != null && !expectedType.isInstance(argument)) {
throw new IllegalArgumentException("Argument %d to function '%s' was %s, expected %s".formatted(i, functionName,
argument.getClass().getName(), expectedType.getName()));
}
}
}
private static void verifyFunctionResult(FunctionCall call, Object result) {
if (result == null) {
return;
}
Class<?> expected = wrap(call.resultType());
if (!expected.isInstance(result)) {
throw new IllegalArgumentException("Function '%s' returned %s, but CriteriaBuilder.function() declared %s"
.formatted(call.name(), result.getClass().getName(), expected.getName()));
}
}
private static Class<?> wrap(Class<?> type) {
if (!type.isPrimitive()) {
return type;
}
if (type == int.class) {
return Integer.class;
}
if (type == long.class) {
return Long.class;
}
if (type == double.class) {
return Double.class;
}
if (type == float.class) {
return Float.class;
}
if (type == short.class) {
return Short.class;
}
if (type == byte.class) {
return Byte.class;
}
if (type == boolean.class) {
return Boolean.class;
}
if (type == char.class) {
return Character.class;
}
if (type == void.class) {
return Void.class;
}
throw new IllegalArgumentException("Unknown primitive type: " + type);
}
// ============================================================
// Proxy Utility for Unimplemented Methods
// ============================================================
private static Object objectMethodOrUnsupported(Object proxy, Method method, Object[] args, String description) {
if (method.getName().equals("toString") && method.getParameterCount() == 0) {
return description;
} else if (method.getName().equals("hashCode") && method.getParameterCount() == 0) {
return System.identityHashCode(proxy);
} else if (method.getName().equals("equals") && method.getParameterCount() == 1) {
return proxy.equals(args[0]);
} else {
throw new UnsupportedOperationException("Unsupported Criteria API operation: " + method);
}
}
}