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