ThreadTxContext.java

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

import io.kaumei.jdbc.tx.KaumeiTxDefinition;
import io.kaumei.jdbc.tx.KaumeiTxException;
import io.kaumei.jdbc.tx.KaumeiTxSynchronization;
import org.jspecify.annotations.Nullable;

import java.sql.Connection;
import java.sql.SQLException;
import java.util.ArrayList;
import java.util.List;

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

class ThreadTxContext {

    private enum TxState {
        RUNNING,
        COMPLETING,
        FINISHED
    }

    private final Connection con;
    private final Connection protectedCon;
    private final KaumeiTxDefinition definition;
    private final List<KaumeiTxSynchronization> synchronizations = new ArrayList<>();

    private @Nullable Boolean autoCommit;
    private @Nullable Integer originalIsolation;
    private @Nullable Boolean originalReadOnly;
    private boolean rollbackOnly;
    private TxState state = TxState.RUNNING;
    private KaumeiTxSynchronization.AfterStatus afterStatus = KaumeiTxSynchronization.AfterStatus.UNKNOWN;
    private @Nullable Throwable terminalFailure;

    ThreadTxContext(Connection con, KaumeiTxDefinition definition) {
        this.con = con;
        this.protectedCon = new ProtectedConnection(con);
        this.definition = definition;
    }

    Connection protectedCon() {
        this.requireNotFinished();
        return this.protectedCon;
    }

    void prepareConnectionState() throws SQLException {
        var autoCommit = con.getAutoCommit();
        if (!autoCommit) {
            // fresh connections must have always the auto commit flag set to true
            throw new KaumeiTxException("Auto commit already disabled.");
        }
        this.autoCommit = autoCommit;
        con.setAutoCommit(false);
        if (this.definition.isolation() != KaumeiTxDefinition.Isolation.DEFAULT) {
            this.originalIsolation = this.con.getTransactionIsolation();
            this.con.setTransactionIsolation(jdbcIsolation(this.definition.isolation()));
        }
        if (this.definition.readOnly() != KaumeiTxDefinition.ReadOnly.DEFAULT) {
            this.originalReadOnly = this.con.isReadOnly();
            this.con.setReadOnly(this.definition.readOnly() == KaumeiTxDefinition.ReadOnly.READ_ONLY);
        }
    }

    private static int jdbcIsolation(KaumeiTxDefinition.Isolation isolation) {
        return switch (isolation) {
            case READ_UNCOMMITTED -> Connection.TRANSACTION_READ_UNCOMMITTED;
            case READ_COMMITTED -> Connection.TRANSACTION_READ_COMMITTED;
            case REPEATABLE_READ -> Connection.TRANSACTION_REPEATABLE_READ;
            case SERIALIZABLE -> Connection.TRANSACTION_SERIALIZABLE;
            case DEFAULT ->
                    throw new IllegalArgumentException("DEFAULT has no JDBC isolation value.");
        };
    }

    @Nullable Throwable restoreConnectionState(@Nullable Throwable exp) {
        if (this.originalReadOnly != null) {
            try {
                this.con.setReadOnly(this.originalReadOnly);
            } catch (Throwable e) {
                exp = addSuppressed(exp, e);
            }
        }
        if (this.originalIsolation != null) {
            try {
                this.con.setTransactionIsolation(this.originalIsolation);
            } catch (Throwable e) {
                exp = addSuppressed(exp, e);
            }
        }
        if (this.autoCommit != null) {
            try {
                this.con.setAutoCommit(this.autoCommit);
            } catch (Throwable e) {
                exp = addSuppressed(exp, e);
            }
        }
        return exp;
    }

    @Nullable Throwable closeContext(@Nullable Throwable exp) {
        // ----- restore read-only, isolation level, auto commit
        exp = this.restoreConnectionState(exp);
        // ----- close
        return close(this.con, exp);
    }

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

    void validateCompatible(KaumeiTxDefinition joined) {
        if (state != TxState.RUNNING) {
            throw new IllegalStateException("Transaction not running.");
        } else if (joined.isolation() != KaumeiTxDefinition.Isolation.DEFAULT
                && joined.isolation() != definition.isolation()) {
            throw new IllegalStateException("Incompatible transaction isolation.");
        } else if (joined.readOnly() != KaumeiTxDefinition.ReadOnly.DEFAULT
                && joined.readOnly() != definition.readOnly()) {
            throw new IllegalStateException("Incompatible transaction read-only.");
        }
    }

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

    boolean isTransactionFinished() {
        return this.state == TxState.FINISHED;
    }

    boolean isTransactionActive() {
        return this.state != TxState.FINISHED;
    }

    void completing() {
        if (state != TxState.RUNNING) {
            throw new IllegalStateException(
                    "Invalid transaction state transition from " + state + " to " + TxState.COMPLETING + ".");
        }
        state = TxState.COMPLETING;
    }

    void running() {
        if (state != TxState.COMPLETING) {
            throw new IllegalStateException(
                    "Invalid transaction state transition from " + state + " to " + TxState.RUNNING + ".");
        }
        this.state = TxState.RUNNING;
        afterStatus = KaumeiTxSynchronization.AfterStatus.UNKNOWN;
    }

    void finished(@Nullable Throwable failure) {
        if (state != TxState.COMPLETING) {
            throw new IllegalStateException(
                    "Invalid transaction state transition from " + state + " to " + TxState.FINISHED + ".");
        }
        terminalFailure = failure;
        state = TxState.FINISHED;
    }

    Throwable terminalFailure(@Nullable Throwable exp) {
        Throwable failure = requireNonNull(terminalFailure, "terminalFailure");
        if (exp != null && exp != failure) {
            failure.addSuppressed(exp);
        }
        return failure;
    }

    void requireNotFinished() {
        if (state == TxState.FINISHED) {
            throw new IllegalStateException("Transaction already completed.");
        }
    }

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

    boolean isRollbackOnly() {
        this.requireNotFinished();
        return this.rollbackOnly;
    }

    void setRollbackOnly() {
        this.requireNotFinished();
        this.rollbackOnly = true;
    }

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

    void addSynchronization(KaumeiTxSynchronization synchronization) {
        if (this.state != TxState.RUNNING) {
            throw new IllegalStateException("Transaction completion already started.");
        }
        this.synchronizations.add(synchronization);
    }

    @Nullable Throwable beforeCompletion(@Nullable Throwable exp) {
        if (!this.synchronizations.isEmpty()) {
            var status = this.rollbackOnly
                    ? KaumeiTxSynchronization.BeforeStatus.ROLLING_BACK
                    : KaumeiTxSynchronization.BeforeStatus.COMMITTING;
            for (var synchronization : this.synchronizations) {
                try {
                    var next = requireNonNull(synchronization.beforeCompletion(status));
                    if (next == KaumeiTxSynchronization.BeforeStatus.ROLLING_BACK) {
                        this.rollbackOnly = true;
                        status = KaumeiTxSynchronization.BeforeStatus.ROLLING_BACK;
                    }
                } catch (Throwable e) {
                    this.rollbackOnly = true;
                    status = KaumeiTxSynchronization.BeforeStatus.ROLLING_BACK;
                    exp = addSuppressed(exp, e);
                }
            }
        }
        return exp;
    }

    @Nullable Throwable completeCurrentTransaction(@Nullable Throwable exp) {
        if (this.rollbackOnly) {
            exp = this.rollback(exp);
        } else {
            exp = this.commit(exp);
        }
        return this.afterCompletion(exp);
    }

    private Throwable rollback(@Nullable Throwable exp) {
        try {
            this.con.rollback();
            this.afterStatus = KaumeiTxSynchronization.AfterStatus.ROLLED_BACK;
            return exp != null ? exp : new KaumeiTxException("Transaction rolled back.");
        } catch (Throwable e) {
            this.afterStatus = KaumeiTxSynchronization.AfterStatus.UNKNOWN;
            return addSuppressed(exp, e);
        }
    }

    private @Nullable Throwable commit(@Nullable Throwable exp) {
        try {
            this.con.commit();
            this.afterStatus = KaumeiTxSynchronization.AfterStatus.COMMITTED;
            return exp;
        } catch (Throwable e) {
            this.afterStatus = KaumeiTxSynchronization.AfterStatus.UNKNOWN;
            exp = addSuppressed(exp, e);
            try {
                this.con.rollback();
            } catch (Throwable ee) {
                exp = addSuppressed(exp, ee);
            }
            return exp;
        }
    }

    private @Nullable Throwable afterCompletion(@Nullable Throwable exp) {
        if (!this.synchronizations.isEmpty()) {
            for (KaumeiTxSynchronization synchronization : this.synchronizations) {
                try {
                    synchronization.afterCompletion(this.afterStatus);
                } catch (Throwable e) {
                    exp = addSuppressed(exp, e);
                }
            }
            this.synchronizations.clear();
        }
        return exp;
    }

}