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);
        }
    }
}