InterfaceWrapper.java

/*
 * SPDX-FileCopyrightText: 2026 kaumei.io
 * SPDX-License-Identifier: Apache-2.0
 */
package io.kaumei.jdbc.tx.internal;

import org.jspecify.annotations.Nullable;

import java.lang.reflect.*;
import java.util.HashMap;
import java.util.Map;

import static java.util.Objects.requireNonNull;

/**
 * Wraps a public interface around an arbitrary target object.
 * The target does not need to implement the interface; it only needs matching public methods.
 * This is primarily intended as small test glue for adding behaviour around calls to an existing
 * class, including a final class.
 *
 * <p>For example:
 * <pre>{@code
 * public interface Calculator {
 *     int add(int left, int right);
 * }
 *
 * final class ExistingCalculator {
 *     public int add(int left, int right) {
 *         return left + right;
 *     }
 * }
 *
 * Calculator calculator = InterfaceWrapper.wrap(
 *         Calculator.class,
 *         new ExistingCalculator(),
 *         invocation -> {
 *             System.out.println(invocation.method().getName());
 *             return invocation.proceed();
 *         });
 *
 * int result = calculator.add(2, 3);
 * }</pre>
 *
 * <p>Only abstract interface methods are forwarded through the interceptor.
 * Default interface methods execute on the wrapper and are not forwarded or directly
 * intercepted.
 */
public final class InterfaceWrapper {

    private static final @Nullable Object[] NO_ARGUMENTS = new Object[0];

    private InterfaceWrapper() {
    }

    public static <T> T wrap(Class<T> api, Object target, Interceptor interceptor) {
        requireNonNull(api, "api");
        requireNonNull(target, "target");
        requireNonNull(interceptor, "interceptor");

        if (!api.isInterface()) {
            throw new IllegalArgumentException("API must be an interface: " + api.getName());
        }
        if (!Modifier.isPublic(api.getModifiers())) {
            throw new IllegalArgumentException("API must be public: " + api.getName());
        }

        Map<Method, Method> methods = resolveMethods(api, target.getClass());
        InvocationHandler handler = new InterceptorInvocationHandler(target, interceptor, methods);
        Object proxy = Proxy.newProxyInstance(api.getClassLoader(), new Class<?>[]{api}, handler);
        return api.cast(proxy);
    }

    @FunctionalInterface
    public interface Interceptor {

        @Nullable Object intercept(Invocation invocation) throws Throwable;
    }

    public interface Invocation {

        /**
         * Returns the invoked interface method.
         */
        Method method();

        @Nullable Object[] arguments();

        /**
         * Invokes the matching method on the target object.
         */
        @Nullable Object proceed() throws Throwable;
    }

    private static final class InterceptorInvocationHandler implements InvocationHandler {

        private final Object target;
        private final Interceptor interceptor;
        private final Map<Method, Method> methods;

        private InterceptorInvocationHandler(
                Object target,
                Interceptor interceptor,
                Map<Method, Method> methods) {
            this.target = target;
            this.interceptor = interceptor;
            this.methods = methods;
        }

        @Override
        public @Nullable Object invoke(
                Object proxy,
                Method method,
                @Nullable Object @Nullable [] arguments) throws Throwable {
            @Nullable Object[] invocationArguments = arguments == null ? NO_ARGUMENTS : arguments;

            if (method.getDeclaringClass() == Object.class) {
                return invokeObjectMethod(proxy, method, invocationArguments);
            }
            if (method.isDefault()) {
                return InvocationHandler.invokeDefault(proxy, method, invocationArguments);
            }
            if (!Modifier.isAbstract(method.getModifiers())) {
                throw new UnsupportedOperationException(method.toString());
            }

            Method targetMethod = methods.get(method);
            if (targetMethod == null) {
                throw new IllegalStateException("No resolved target method for " + method);
            }
            return interceptor.intercept(
                    new MethodInvocation(method, target, targetMethod, invocationArguments));
        }

        private Object invokeObjectMethod(Object proxy, Method method, @Nullable Object[] arguments) {
            return switch (method.getName()) {
                case "toString" -> "InterfaceWrapper[" + target + "]";
                case "hashCode" -> System.identityHashCode(proxy);
                case "equals" -> proxy == arguments[0];
                default -> throw new UnsupportedOperationException(method.toString());
            };
        }
    }

    private static final class MethodInvocation implements Invocation {

        private final Object target;
        private final Method interfaceMethod;
        private final Method targetMethod;
        private final @Nullable Object[] arguments;

        private MethodInvocation(
                Method interfaceMethod,
                Object target,
                Method targetMethod,
                @Nullable Object[] arguments) {
            this.interfaceMethod = interfaceMethod;
            this.target = target;
            this.targetMethod = targetMethod;
            this.arguments = arguments;
        }

        @Override
        public Method method() {
            return interfaceMethod;
        }

        @Override
        public @Nullable Object[] arguments() {
            return arguments;
        }

        @Override
        public @Nullable Object proceed() throws Throwable {
            try {
                return targetMethod.invoke(target, arguments);
            } catch (InvocationTargetException e) {
                throw e.getTargetException();
            }
        }
    }

    private static Map<Method, Method> resolveMethods(Class<?> api, Class<?> targetType) {
        Map<Method, Method> methods = new HashMap<>();
        for (Method interfaceMethod : api.getMethods()) {
            if (!Modifier.isAbstract(interfaceMethod.getModifiers())) {
                continue;
            }

            Method targetMethod;
            try {
                targetMethod = targetType.getMethod(
                        interfaceMethod.getName(),
                        interfaceMethod.getParameterTypes());
            } catch (NoSuchMethodException e) {
                throw new IllegalArgumentException(
                        "No matching public method on " + targetType.getName()
                                + " for " + methodSignature(interfaceMethod),
                        e);
            }

            if (!interfaceMethod.getReturnType().isAssignableFrom(targetMethod.getReturnType())) {
                throw new IllegalArgumentException(
                        "Incompatible return type for " + methodSignature(interfaceMethod));
            }
            if (!targetMethod.trySetAccessible()) {
                throw new IllegalArgumentException(
                        "Target method is not accessible: " + targetMethod);
            }
            methods.put(interfaceMethod, targetMethod);
        }
        return Map.copyOf(methods);
    }

    private static String methodSignature(Method method) {
        StringBuilder signature = new StringBuilder();
        signature.append(method.getName());
        signature.append('(');
        Class<?>[] parameterTypes = method.getParameterTypes();
        for (int i = 0; i < parameterTypes.length; i++) {
            if (i > 0) {
                signature.append(", ");
            }
            signature.append(parameterTypes[i].getSimpleName());
        }
        signature.append(')');
        return signature.toString();
    }
}