KaumeiTxManagerImpl.java

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

import io.kaumei.jdbc.core.JdbcConnectionProvider;
import io.kaumei.jdbc.tx.*;
import org.jspecify.annotations.Nullable;

import java.sql.Connection;
import java.sql.SQLException;

import static io.kaumei.jdbc.tx.internal.Utils.close;
import static io.kaumei.jdbc.tx.internal.Utils.throwUnchecked;
import static java.util.Objects.requireNonNull;

public class KaumeiTxManagerImpl implements KaumeiTxManager {

    private final JdbcConnectionProvider provider;
    private final ActiveTxContext activeTx = new ActiveTxContext();

    public KaumeiTxManagerImpl(JdbcConnectionProvider provider) {
        this.provider = requireNonNull(provider, "provider");
    }

    // ------------------------------------------------------------------------

    @Override
    public Connection getConnection() throws SQLException {
        var ctx = activeTx.currentOpt();
        if (ctx == null) {
            return provider.getConnection();
        }
        return ctx.protectedCon();
    }

    // ------------------------------------------------------------------------

    @Override
    public void setRollbackOnly() {
        activeTx.current().setRollbackOnly();
    }

    @Override
    public boolean isRollbackOnly() {
        return activeTx.current().isRollbackOnly();
    }

    @Override
    public boolean isTransactionActive() {
        var ctx = activeTx.currentOpt();
        return ctx != null && ctx.isTransactionActive();
    }

    @Override
    public void registerSynchronization(KaumeiTxSynchronization synchronization) {
        requireNonNull(synchronization, "synchronization");
        activeTx.current().addSynchronization(synchronization);
    }

    @Override
    public void commitAndContinue() {
        ThreadTxContext ctx = activeTx.current();
        ctx.completing();
        Throwable exp = ctx.beforeCompletion(null);
        activeTx.suspend();
        try {
            exp = ctx.completeCurrentTransaction(exp);
        } finally {
            activeTx.restore(ctx);
        }
        if (exp == null) {
            ctx.running();
        } else {
            ctx.finished(exp);
            throwUnchecked(exp);
        }
    }

    @Override
    public void required(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.REQUIRED, definition, callback);
    }

    @Override
    public <T> T required(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.REQUIRED, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T requiredOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.REQUIRED, definition, callback);
    }

    @Override
    public void requiresNew(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.REQUIRES_NEW, definition, callback);
    }

    @Override
    public <T> T requiresNew(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.REQUIRES_NEW, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T requiresNewOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.REQUIRES_NEW, definition, callback);
    }

    @Override
    public void mandatory(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.MANDATORY, definition, callback);
    }

    @Override
    public <T> T mandatory(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.MANDATORY, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T mandatoryOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.MANDATORY, definition, callback);
    }

    @Override
    public void supports(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.SUPPORTS, definition, callback);
    }

    @Override
    public <T> T supports(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.SUPPORTS, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T supportsOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.SUPPORTS, definition, callback);
    }

    @Override
    public void notSupported(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.NOT_SUPPORTED, definition, callback);
    }

    @Override
    public <T> T notSupported(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.NOT_SUPPORTED, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T notSupportedOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.NOT_SUPPORTED, definition, callback);
    }

    @Override
    public void never(KaumeiTxDefinition definition, Callback callback) throws Exception {
        execute(KaumeiTxType.NEVER, definition, callback);
    }

    @Override
    public <T> T never(KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        var result = execute(KaumeiTxType.NEVER, definition, callback);
        return requireNonNull(result, "result"); // sanity check
    }

    @Override
    public <T> @Nullable T neverOpt(KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        return executeOpt(KaumeiTxType.NEVER, definition, callback);
    }

    // ------------------------------------------------------------------------

    private void execute(KaumeiTxType type, KaumeiTxDefinition definition, Callback callback) throws Exception {
        requireNonNull(type, "type");
        requireNonNull(definition, "definition");
        requireNonNull(callback, "callback");
        Throwable exp = null;
        var scope = before(type, definition);
        var context = new CallbackContext(scope.ctx);
        try {
            callback.execute(context);
        } catch (Throwable e) {
            exp = exception(type, definition, scope, e);
        } finally {
            exp = after(type, scope, exp);
        }

        if (exp instanceof Exception e) {
            throw e;
        } else if (exp instanceof Error e) {
            throw e;
        } else if (exp != null) {
            throw new KaumeiTxException(exp.getMessage(), exp);
        }
    }

    private <T> @Nullable T execute(KaumeiTxType type, KaumeiTxDefinition definition, CallbackWithResult<T> callback) throws Exception {
        requireNonNull(type, "type");
        requireNonNull(definition, "definition");
        requireNonNull(callback, "callback");
        Throwable exp = null;
        T result = null;
        var scope = before(type, definition);
        var context = new CallbackContext(scope.ctx);
        try {
            result = callback.execute(context);
        } catch (Throwable e) {
            exp = exception(type, definition, scope, e);
        } finally {
            exp = after(type, scope, exp);
        }

        if (exp instanceof Exception e) {
            throw e;
        } else if (exp instanceof Error e) {
            throw e;
        } else if (exp != null) {
            throw new KaumeiTxException(exp.getMessage(), exp);
        }
        return result;
    }

    private <T> @Nullable T executeOpt(KaumeiTxType type, KaumeiTxDefinition definition, CallbackWithNullableResult<T> callback) throws Exception {
        requireNonNull(type, "type");
        requireNonNull(definition, "definition");
        requireNonNull(callback, "callback");
        Throwable exp = null;
        T result = null;
        var scope = before(type, definition);
        var context = new CallbackContext(scope.ctx);
        try {
            result = callback.execute(context);
        } catch (Throwable e) {
            exp = exception(type, definition, scope, e);
        } finally {
            exp = after(type, scope, exp);
        }

        if (exp instanceof Exception e) {
            throw e;
        } else if (exp instanceof Error e) {
            throw e;
        } else if (exp != null) {
            throw new KaumeiTxException(exp.getMessage(), exp);
        }

        return result;
    }

    // ------------------------------------------------------------------------

    private TxScope before(KaumeiTxType type, KaumeiTxDefinition definition) {
        var current = activeTx.currentOpt();
        if (current != null) {
            current.requireNotFinished();
        }
        return switch (type) {
            case REQUIRED -> activeTx.hasCurrent()
                    ? join(definition)
                    : start(definition, null, false);
            case REQUIRES_NEW -> {
                var suspended = activeTx.suspend();
                try {
                    yield start(definition, suspended, true);
                } catch (RuntimeException | Error e) {
                    activeTx.restore(suspended);
                    throw e;
                }
            }
            case MANDATORY -> {
                activeTx.current();
                yield join(definition);
            }
            case SUPPORTS -> activeTx.hasCurrent()
                    ? join(definition)
                    : TxScope.none();
            case NOT_SUPPORTED -> TxScope.suspended(activeTx.suspend());
            case NEVER -> {
                activeTx.requireNoCurrent();
                yield TxScope.none();
            }
        };
    }

    private TxScope join(KaumeiTxDefinition definition) {
        var ctx = activeTx.current();
        ctx.validateCompatible(definition);
        return TxScope.join(ctx);
    }

    private TxScope start(KaumeiTxDefinition definition, @Nullable ThreadTxContext suspended, boolean restoreSuspended) {
        ThreadTxContext ctx = startTx(definition);
        activeTx.activate(ctx);
        return TxScope.started(ctx, suspended, restoreSuspended);
    }

    // ------------------------------------------------------------------------

    private Throwable exception(KaumeiTxType type, KaumeiTxDefinition definition, TxScope scope, Throwable exp) {
        return switch (type) {
            case REQUIRED, MANDATORY, REQUIRES_NEW ->
                    markRollbackOnly(requireNonNull(scope.ctx, "ctx"), definition, exp);
            case SUPPORTS -> scope.ctx == null
                    ? exp
                    : markRollbackOnly(scope.ctx, definition, exp);
            case NOT_SUPPORTED, NEVER -> exp;
        };
    }

    Throwable markRollbackOnly(ThreadTxContext ctx, KaumeiTxDefinition definition, Throwable exp) {
        if (ctx.isTransactionFinished() || ctx.isRollbackOnly()) {
            return exp;
        } else if (exp instanceof Error) {
            ctx.setRollbackOnly();
        } else {
            try {
                if (definition.dontRollbackOn().matches(exp)) {
                    return exp;
                } else if (definition.rollbackOn().matches(exp)
                        || exp instanceof RuntimeException) {
                    ctx.setRollbackOnly();
                }
            } catch (Throwable matcherFailure) {
                ctx.setRollbackOnly();
                if (matcherFailure != exp) {
                    exp.addSuppressed(matcherFailure);
                }
            }
        }
        return exp;
    }

    // ------------------------------------------------------------------------

    private @Nullable Throwable after(KaumeiTxType type, TxScope scope, @Nullable Throwable exp) {
        return switch (type) {
            case REQUIRED -> scope.completesTx
                    ? completeTx(requireNonNull(scope.ctx, "ctx"), exp)
                    : exp;
            case REQUIRES_NEW -> {
                try {
                    yield completeTx(requireNonNull(scope.ctx, "ctx"), exp);
                } finally {
                    activeTx.restore(scope.suspended);
                }
            }
            case MANDATORY, SUPPORTS, NEVER -> exp;
            case NOT_SUPPORTED -> {
                activeTx.restore(scope.suspended);
                yield exp;
            }
        };
    }

    private @Nullable Throwable completeTx(ThreadTxContext ctx, @Nullable Throwable exp) {
        if (ctx.isTransactionFinished()) {
            activeTx.suspend();
            return ctx.closeContext(ctx.terminalFailure(exp));
        }

        try {
            ctx.completing();
            exp = ctx.beforeCompletion(exp);
        } finally {
            activeTx.suspend();
        }

        exp = ctx.completeCurrentTransaction(exp);
        ctx.finished(exp);
        return ctx.closeContext(exp);
    }

    // ------------------------------------------------------------------------

    private ThreadTxContext startTx(KaumeiTxDefinition definition) {
        Connection con = null;
        ThreadTxContext ctx = null;
        try {
            con = provider.getConnection();
            ctx = new ThreadTxContext(con, definition);
            ctx.prepareConnectionState();
            return ctx;
        } catch (Throwable e) {
            Throwable exp = e;
            if (ctx != null) {
                exp = ctx.restoreConnectionState(exp);
            }
            if (con != null) {
                exp = close(con, exp);
            }
            if (exp instanceof RuntimeException re) {
                throw re;
            } else if (exp instanceof Error error) {
                throw error;
            } else {
                throw new KaumeiTxException(exp.getMessage(), exp);
            }
        }
    }

    private final class CallbackContext implements KaumeiTxContext {

        private final @Nullable ThreadTxContext expected;

        private CallbackContext(@Nullable ThreadTxContext expected) {
            this.expected = expected;
        }

        @Override
        public Connection getConnection() throws SQLException {
            KaumeiTxManagerImpl txManager = this.txManager();
            if (expected == null) {
                throw new IllegalStateException("No active transaction.");
            }
            return txManager.getConnection();
        }

        @Override
        public boolean isTransactionActive() {
            return this.txManager().isTransactionActive();
        }

        @Override
        public void setRollbackOnly() {
            this.txManager().setRollbackOnly();
        }

        @Override
        public boolean isRollbackOnly() {
            return this.txManager().isRollbackOnly();
        }

        @Override
        public void registerSynchronization(KaumeiTxSynchronization synchronization) {
            this.txManager().registerSynchronization(synchronization);
        }

        @Override
        public void commitAndContinue() {
            this.txManager().commitAndContinue();
        }

        private KaumeiTxManagerImpl txManager() {
            if (activeTx.currentOpt() != expected) {
                throw new IllegalStateException("Transaction callback context is not current.");
            }
            return KaumeiTxManagerImpl.this;
        }
    }

    // ------------------------------------------------------------------------
    // static stuff

    private record TxScope(@Nullable ThreadTxContext ctx, @Nullable ThreadTxContext suspended,
                           boolean completesTx, boolean restoreSuspended) {

        static TxScope none() {
            return new TxScope(null, null, false, false);
        }

        static TxScope join(ThreadTxContext ctx) {
            return new TxScope(ctx, null, false, false);
        }

        static TxScope started(ThreadTxContext ctx, @Nullable ThreadTxContext suspended, boolean restoreSuspended) {
            return new TxScope(ctx, suspended, true, restoreSuspended);
        }

        static TxScope suspended(@Nullable ThreadTxContext suspended) {
            return new TxScope(null, suspended, false, true);
        }
    }

}