KaumeiTxProxy.java

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

import io.kaumei.jdbc.annotation.KaumeiTx;
import org.jspecify.annotations.Nullable;

import java.lang.invoke.MethodHandle;
import java.lang.invoke.MethodHandles;
import java.lang.invoke.MethodType;
import java.lang.reflect.InvocationHandler;
import java.lang.reflect.Method;
import java.lang.reflect.Modifier;
import java.lang.reflect.Proxy;
import java.util.*;

import static java.util.Objects.requireNonNull;

/**
 * Wraps an annotated interface around an implementation and applies its {@link KaumeiTx}
 * declarations through a {@link KaumeiTxManager}.
 *
 * <p>Only annotations on the interface and its methods are considered. Annotation resolution and
 * transaction definitions are prepared when the proxy is created. The target must obtain its JDBC
 * connections from the same transaction manager so its database work participates in these
 * transactions.
 *
 * <p>Void methods use the transaction manager's void callback overload. Methods with a return
 * value use its nullable result overload; the proxy does not enforce business return contracts.
 */
public final class KaumeiTxProxy {

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

    private final Object target;
    private final KaumeiTxManager txManager;
    private final Map<Method, MethodPlan> plans;

    private KaumeiTxProxy(Object target, KaumeiTxManager txManager, Map<Method, MethodPlan> plans) {
        this.target = target;
        this.txManager = txManager;
        this.plans = plans;
    }

    /**
     * Creates a transaction-aware proxy for {@code target}.
     * @param api interface exposed by the proxy
     * @param target implementation of {@code api}
     * @param txManager transaction manager used for annotated calls and by the target as its
     * {@code JdbcConnectionProvider}
     */
    public static <T> T wrap(Class<T> api, T target, KaumeiTxManager txManager) {
        requireNonNull(api, "api");
        requireNonNull(target, "target");
        requireNonNull(txManager, "txManager");

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

        List<List<Class<?>>> interfaceLevels = interfaceLevels(api);
        Map<Method, MethodPlan> plans = new LinkedHashMap<>();
        for (Method method : api.getMethods()) {
            if (Modifier.isStatic(method.getModifiers())) {
                continue;
            }
            KaumeiTx annotation = resolveAnnotation(method, interfaceLevels);
            boolean invokeDefault = invokesInterfaceDefault(method, target.getClass());
            plans.put(method, methodPlan(method, target, annotation, invokeDefault));
        }

        KaumeiTxProxy handler = new KaumeiTxProxy(target, txManager, Map.copyOf(plans));
        Object proxy = Proxy.newProxyInstance(api.getClassLoader(), new Class<?>[]{api}, handler::invokeTransactional);
        return api.cast(proxy);
    }

    private static MethodPlan methodPlan(Method method, Object target, @Nullable KaumeiTx annotation,
                                         boolean invokeDefault) {
        ReturnModel returnModel = method.getReturnType() == void.class ? ReturnModel.VOID : ReturnModel.VALUE;
        MethodHandle targetMethod = invokeDefault ? null : targetMethod(method, target);
        if (annotation == null) {
            return new MethodPlan(null, KaumeiTxDefinition.DEFAULT, returnModel, method, targetMethod,
                    invokeDefault);
        }
        return new MethodPlan(
                annotation.value(),
                KaumeiTxDefinition.from(annotation),
                returnModel,
                method,
                targetMethod,
                invokeDefault);
    }

    private static MethodHandle targetMethod(Method method, Object target) {
        try {
            MethodHandle targetMethod = MethodHandles.publicLookup()
                    .unreflect(method)
                    .bindTo(target)
                    .asSpreader(Object[].class, method.getParameterCount());
            return targetMethod.asType(MethodType.methodType(Object.class, Object[].class));
        } catch (IllegalAccessException e) {
            throw new IllegalArgumentException("Cannot access target method " + method.toGenericString(), e);
        }
    }

    private static @Nullable KaumeiTx resolveAnnotation(Method method, List<List<Class<?>>> interfaceLevels) {
        for (List<Class<?>> level : interfaceLevels) {
            List<AnnotationSource> annotations = new ArrayList<>();
            boolean declarationFound = false;
            for (Class<?> candidate : level) {
                Method declaration = declaredMethod(candidate, method);
                if (declaration == null) {
                    continue;
                }
                declarationFound = true;
                KaumeiTx annotation = declaration.getDeclaredAnnotation(KaumeiTx.class);
                if (annotation != null) {
                    annotations.add(new AnnotationSource(annotation, declaration.toGenericString()));
                }
            }
            if (declarationFound) {
                if (!annotations.isEmpty()) {
                    return uniqueAnnotation(method, annotations);
                }
                break;
            }
        }

        for (List<Class<?>> level : interfaceLevels) {
            List<AnnotationSource> annotations = new ArrayList<>();
            for (Class<?> candidate : level) {
                KaumeiTx annotation = candidate.getDeclaredAnnotation(KaumeiTx.class);
                if (annotation != null) {
                    annotations.add(new AnnotationSource(annotation, candidate.getName()));
                }
            }
            if (!annotations.isEmpty()) {
                return uniqueAnnotation(method, annotations);
            }
        }
        return null;
    }

    private static @Nullable Method declaredMethod(Class<?> candidate, Method method) {
        try {
            return candidate.getDeclaredMethod(method.getName(), method.getParameterTypes());
        } catch (NoSuchMethodException e) {
            return null;
        }
    }

    private static KaumeiTx uniqueAnnotation(Method method, List<AnnotationSource> sources) {
        KaumeiTx selected = sources.get(0).annotation();
        for (int i = 1; i < sources.size(); i++) {
            if (!selected.equals(sources.get(i).annotation())) {
                StringBuilder message = new StringBuilder();
                message.append("Conflicting @KaumeiTx annotations for ");
                message.append(method.toGenericString());
                message.append(" on ");
                for (int sourceIndex = 0; sourceIndex < sources.size(); sourceIndex++) {
                    if (sourceIndex > 0) {
                        message.append(", ");
                    }
                    message.append(sources.get(sourceIndex).source());
                }
                throw new IllegalArgumentException(message.toString());
            }
        }
        return selected;
    }

    private static List<List<Class<?>>> interfaceLevels(Class<?> api) {
        List<List<Class<?>>> levels = new ArrayList<>();
        Set<Class<?>> visited = new HashSet<>();
        List<Class<?>> current = List.of(api);
        while (!current.isEmpty()) {
            List<Class<?>> level = new ArrayList<>();
            List<Class<?>> next = new ArrayList<>();
            for (Class<?> candidate : current) {
                if (!visited.add(candidate)) {
                    continue;
                }
                level.add(candidate);
                for (Class<?> parent : candidate.getInterfaces()) {
                    if (!visited.contains(parent)) {
                        next.add(parent);
                    }
                }
            }
            level.sort(Comparator.comparing(Class::getName));
            if (!level.isEmpty()) {
                levels.add(List.copyOf(level));
            }
            current = next;
        }
        return List.copyOf(levels);
    }

    private static boolean invokesInterfaceDefault(Method method, Class<?> targetType) {
        if (!method.isDefault()) {
            return false;
        }
        try {
            Method targetMethod = targetType.getMethod(method.getName(), method.getParameterTypes());
            return targetMethod.getDeclaringClass().isInterface();
        } catch (NoSuchMethodException e) {
            throw new IllegalArgumentException("No matching target method for " + method.toGenericString(), e);
        }
    }

    private enum ReturnModel {
        VOID,
        VALUE
    }

    private record MethodPlan(KaumeiTx.@Nullable Type type, KaumeiTxDefinition definition,
                              ReturnModel returnModel, Method method,
                              @Nullable MethodHandle targetMethod, boolean invokeDefault) {
    }

    private record AnnotationSource(KaumeiTx annotation, String source) {
    }

    private @Nullable Object invokeTransactional(
            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);
        }

        MethodPlan plan = plans.get(method);
        if (plan == null) {
            throw new IllegalStateException("No transaction plan for " + method.toGenericString());
        }
        KaumeiTxBusinessMethod businessMethod = new KaumeiTxBusinessMethod(
                proxy,
                plan.method(),
                plan.targetMethod(),
                plan.invokeDefault(),
                invocationArguments);
        if (plan.type() == null) {
            return businessMethod.invoke(null);
        }
        if (plan.returnModel() == ReturnModel.VOID) {
            KaumeiTxManager.Callback callback = businessMethod::invoke;
            switch (plan.type()) {
                case REQUIRED -> txManager.required(plan.definition(), callback);
                case REQUIRES_NEW -> txManager.requiresNew(plan.definition(), callback);
                case MANDATORY -> txManager.mandatory(plan.definition(), callback);
                case SUPPORTS -> txManager.supports(plan.definition(), callback);
                case NOT_SUPPORTED -> txManager.notSupported(plan.definition(), callback);
                case NEVER -> txManager.never(plan.definition(), callback);
            }
            return null;
        }

        KaumeiTxManager.CallbackWithNullableResult<Object> callback = businessMethod::invoke;
        return switch (plan.type()) {
            case REQUIRED -> txManager.requiredOpt(plan.definition(), callback);
            case REQUIRES_NEW -> txManager.requiresNewOpt(plan.definition(), callback);
            case MANDATORY -> txManager.mandatoryOpt(plan.definition(), callback);
            case SUPPORTS -> txManager.supportsOpt(plan.definition(), callback);
            case NOT_SUPPORTED -> txManager.notSupportedOpt(plan.definition(), callback);
            case NEVER -> txManager.neverOpt(plan.definition(), callback);
        };
    }

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

final class KaumeiTxBusinessMethod {

    private final Object proxy;
    private final Method method;
    private final @Nullable MethodHandle targetMethod;
    private final boolean invokeDefault;
    private final @Nullable Object[] arguments;

    KaumeiTxBusinessMethod(Object proxy, Method method, @Nullable MethodHandle targetMethod,
                           boolean invokeDefault, @Nullable Object[] arguments) {
        this.proxy = proxy;
        this.method = method;
        this.targetMethod = targetMethod;
        this.invokeDefault = invokeDefault;
        this.arguments = arguments;
    }

    @Nullable Object invoke(@Nullable KaumeiTxContext context) {
        try {
            if (invokeDefault) {
                return InvocationHandler.invokeDefault(proxy, method, arguments);
            }
            return (Object) requireNonNull(targetMethod, "targetMethod").invokeExact(arguments);
        } catch (Throwable e) {
            return KaumeiTxBusinessMethod.<Object, RuntimeException>throwUnchecked(e);
        }
    }

    @SuppressWarnings("unchecked")
    private static <T, E extends Throwable> T throwUnchecked(Throwable exception) throws E {
        throw (E) exception;
    }
}