Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -110,9 +110,8 @@ protected void setConnection(@Nullable Connection connection) {

/**
* Return the current {@link Connection} held by this {@link ConnectionHolder}.
* <p>This will be the same {@link Connection} until {@code released} gets called
* on the {@link ConnectionHolder}, which will reset the held {@link Connection},
* fetching a new {@link Connection} on demand.
* <p>This will be the same {@link Connection} until it is reset for a holder
* that is not synchronized with a transaction, fetching a new Connection on demand.
* @see #released()
*/
public Connection getConnection() {
Expand All @@ -137,7 +136,7 @@ String nextSavepoint() {
@Override
public void released() {
super.released();
if (!isOpen() && this.currentConnection != null) {
if (!isOpen() && !isSynchronizedWithTransaction() && this.currentConnection != null) {
this.currentConnection = null;
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,10 @@
package org.springframework.r2dbc.connection;

import java.util.List;
import java.util.concurrent.atomic.AtomicInteger;

import io.r2dbc.spi.Connection;
import io.r2dbc.spi.ConnectionFactory;
import io.r2dbc.spi.R2dbcBadGrammarException;
import io.r2dbc.spi.R2dbcDataIntegrityViolationException;
import io.r2dbc.spi.R2dbcException;
Expand All @@ -30,6 +33,8 @@
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
import org.junit.jupiter.params.provider.FieldSource;
import reactor.core.publisher.Mono;
import reactor.test.StepVerifier;

import org.springframework.dao.CannotAcquireLockException;
import org.springframework.dao.DataAccessResourceFailureException;
Expand All @@ -41,9 +46,16 @@
import org.springframework.dao.TransientDataAccessResourceException;
import org.springframework.r2dbc.BadSqlGrammarException;
import org.springframework.r2dbc.UncategorizedR2dbcException;
import org.springframework.transaction.TransactionDefinition;
import org.springframework.transaction.reactive.TransactionalOperator;
import org.springframework.transaction.support.DefaultTransactionDefinition;

import static org.assertj.core.api.Assertions.assertThat;
import static org.junit.jupiter.params.provider.Arguments.arguments;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;

/**
* Tests for {@link ConnectionFactoryUtils}.
Expand All @@ -53,6 +65,76 @@
*/
class ConnectionFactoryUtilsTests {

@Test
void explicitlyReleasedConnectionIsReusedAndClosedOnTransactionCompletion() {
AtomicInteger transactionConnectionCloses = new AtomicInteger();
TransactionalOperator operator = transactionalOperator(transactionConnectionCloses);
Connection connection = mock();
AtomicInteger connectionCloses = closeCounter(connection);
ConnectionFactory connectionFactory = mock();
when(connectionFactory.create()).thenAnswer(invocation -> Mono.just(connection));

Mono<Connection> useAndRelease = ConnectionFactoryUtils.getConnection(connectionFactory)
.flatMap(con -> ConnectionFactoryUtils.releaseConnection(con, connectionFactory).thenReturn(con));

StepVerifier.create(useAndRelease.then(useAndRelease).as(operator::transactional))
.expectNext(connection)
.verifyComplete();
assertThat(connectionCloses).hasValue(1);
assertThat(transactionConnectionCloses).hasValue(1);
verify(connectionFactory).create();
}

@Test
void connectionReleasedBeforeSuspensionIsClosedAndReplacedAfterResume() {
AtomicInteger transactionConnectionCloses = new AtomicInteger();
TransactionalOperator operator = transactionalOperator(transactionConnectionCloses);
TransactionalOperator requiresNewOperator = transactionalOperator(
transactionConnectionCloses, TransactionDefinition.PROPAGATION_REQUIRES_NEW);
Connection firstConnection = mock();
Connection secondConnection = mock();
AtomicInteger firstConnectionCloses = closeCounter(firstConnection);
AtomicInteger secondConnectionCloses = closeCounter(secondConnection);
ConnectionFactory connectionFactory = connectionFactory(firstConnection, secondConnection);

Mono<Connection> useAndRelease = ConnectionFactoryUtils.getConnection(connectionFactory)
.flatMap(con -> ConnectionFactoryUtils.releaseConnection(con, connectionFactory).thenReturn(con));
Mono<Void> suspension = Mono.<Void>empty().as(requiresNewOperator::transactional);

StepVerifier.create(useAndRelease.then(suspension).then(useAndRelease).as(operator::transactional))
.expectNext(secondConnection)
.verifyComplete();
assertThat(firstConnectionCloses).hasValue(1);
assertThat(secondConnectionCloses).hasValue(1);
}

private static TransactionalOperator transactionalOperator(AtomicInteger connectionCloses) {
return transactionalOperator(connectionCloses, TransactionDefinition.PROPAGATION_REQUIRED);
}

private static TransactionalOperator transactionalOperator(AtomicInteger connectionCloses, int propagationBehavior) {
Connection connection = mock();
when(connection.beginTransaction(any(io.r2dbc.spi.TransactionDefinition.class))).thenReturn(Mono.empty());
when(connection.commitTransaction()).thenReturn(Mono.empty());
when(connection.rollbackTransaction()).thenReturn(Mono.empty());
when(connection.close()).thenReturn(Mono.fromRunnable(connectionCloses::incrementAndGet));
return TransactionalOperator.create(new R2dbcTransactionManager(connectionFactory(connection)),
new DefaultTransactionDefinition(propagationBehavior));
}

private static ConnectionFactory connectionFactory(Connection... connections) {
ConnectionFactory connectionFactory = mock();
AtomicInteger index = new AtomicInteger();
when(connectionFactory.create()).thenAnswer(invocation -> Mono.just(connections[index.getAndIncrement()]));
return connectionFactory;
}

private static AtomicInteger closeCounter(Connection connection) {
AtomicInteger closeCalls = new AtomicInteger();
when(connection.close()).thenReturn(Mono.fromRunnable(closeCalls::incrementAndGet));
return closeCalls;
}

@Test
void shouldTranslateTransientResourceException() {
Exception exception = ConnectionFactoryUtils.convertR2dbcException("", "",
Expand Down