2424import static org .hamcrest .CoreMatchers .nullValue ;
2525import static org .hamcrest .MatcherAssert .assertThat ;
2626import static org .junit .Assert .assertEquals ;
27+ import static org .junit .Assert .assertFalse ;
2728import static org .junit .Assert .assertNotNull ;
29+ import static org .junit .Assert .assertNotSame ;
2830import static org .junit .Assert .assertNull ;
31+ import static org .junit .Assert .assertSame ;
2932import static org .junit .Assert .fail ;
3033import static org .mockito .Mockito .any ;
3134import static org .mockito .Mockito .doThrow ;
6770import java .math .BigDecimal ;
6871import java .util .Arrays ;
6972import java .util .Collections ;
73+ import java .util .concurrent .CountDownLatch ;
7074import org .junit .Test ;
7175import org .junit .runner .RunWith ;
7276import org .junit .runners .JUnit4 ;
@@ -158,12 +162,21 @@ private ReadWriteTransaction createSubject() {
158162 return createSubject (CommitBehavior .SUCCEED , false );
159163 }
160164
165+ private ReadWriteTransaction createSubject (boolean keepTransactionAlive ) {
166+ return createSubject (CommitBehavior .SUCCEED , false , keepTransactionAlive );
167+ }
168+
161169 private ReadWriteTransaction createSubject (CommitBehavior commitBehavior ) {
162- return createSubject (commitBehavior , false );
170+ return createSubject (commitBehavior , false , false );
163171 }
164172
165173 private ReadWriteTransaction createSubject (
166174 final CommitBehavior commitBehavior , boolean withRetry ) {
175+ return createSubject (commitBehavior , withRetry , false );
176+ }
177+
178+ private ReadWriteTransaction createSubject (
179+ final CommitBehavior commitBehavior , boolean withRetry , boolean keepTransactionAlive ) {
167180 DatabaseClient client = mock (DatabaseClient .class );
168181 when (client .transactionManager ())
169182 .thenAnswer (
@@ -179,6 +192,7 @@ private ReadWriteTransaction createSubject(
179192 });
180193 return ReadWriteTransaction .newBuilder ()
181194 .setDatabaseClient (client )
195+ .setKeepTransactionAlive (keepTransactionAlive )
182196 .setRetryAbortsInternally (withRetry )
183197 .setIsolationLevel (IsolationLevel .ISOLATION_LEVEL_UNSPECIFIED )
184198 .setSavepointSupport (SavepointSupport .FAIL_AFTER_ROLLBACK )
@@ -857,6 +871,137 @@ public void testGetCommitResponseAfterCommit() {
857871 assertNotNull (transaction .getCommitResponseOrNull ());
858872 }
859873
874+ @ Test
875+ public void testKeepAliveTaskRemovedFromQueueOnCancel () {
876+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
877+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
878+ when (parsedStatement .isUpdate ()).thenReturn (true );
879+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
880+ when (parsedStatement .getStatement ()).thenReturn (statement );
881+
882+ int initialQueueSize = ReadWriteTransaction .getKeepAliveService ().getQueue ().size ();
883+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
884+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
885+
886+ assertEquals (
887+ initialQueueSize + 1 , ReadWriteTransaction .getKeepAliveService ().getQueue ().size ());
888+
889+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
890+ assertEquals (initialQueueSize , ReadWriteTransaction .getKeepAliveService ().getQueue ().size ());
891+ }
892+
893+ @ Test
894+ public void testKeepAliveWeakReference () {
895+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
896+ ReadWriteTransaction .KeepAliveRunnable runnable =
897+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
898+
899+ assertNotNull (runnable .transactionRef );
900+ assertSame (transaction , runnable .transactionRef .get ());
901+ }
902+
903+ @ Test
904+ public void testKeepAliveRescheduledWhenLockBusy () {
905+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
906+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
907+ when (parsedStatement .isUpdate ()).thenReturn (true );
908+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
909+ when (parsedStatement .getStatement ()).thenReturn (statement );
910+
911+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
912+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
913+
914+ ScheduledFuture <?> future1 = transaction .getKeepAliveFuture ();
915+ assertNotNull (future1 );
916+
917+ CountDownLatch latch = new CountDownLatch (1 );
918+ CountDownLatch lockAcquired = new CountDownLatch (1 );
919+ Thread lockHoldingThread =
920+ new Thread (
921+ () -> {
922+ transaction .abortedLock .lock ();
923+ try {
924+ lockAcquired .countDown ();
925+ latch .await ();
926+ } catch (InterruptedException e ) {
927+ Thread .currentThread ().interrupt ();
928+ } finally {
929+ transaction .abortedLock .unlock ();
930+ }
931+ });
932+ lockHoldingThread .start ();
933+ try {
934+ lockAcquired .await ();
935+ ReadWriteTransaction .KeepAliveRunnable runnable =
936+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
937+ runnable .run ();
938+ } catch (InterruptedException e ) {
939+ Thread .currentThread ().interrupt ();
940+ fail ("Test interrupted" );
941+ } finally {
942+ latch .countDown ();
943+ try {
944+ lockHoldingThread .join ();
945+ } catch (InterruptedException e ) {
946+ Thread .currentThread ().interrupt ();
947+ }
948+ }
949+
950+ ScheduledFuture <?> future2 = transaction .getKeepAliveFuture ();
951+ assertNotNull (future2 );
952+ assertNotSame (future1 , future2 );
953+ }
954+
955+ @ Test
956+ public void testKeepAliveFutureNullifiedOnCancel () {
957+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
958+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
959+ when (parsedStatement .isUpdate ()).thenReturn (true );
960+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
961+ when (parsedStatement .getStatement ()).thenReturn (statement );
962+
963+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
964+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
965+
966+ assertNotNull (transaction .getKeepAliveFuture ());
967+
968+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
969+
970+ assertNull (transaction .getKeepAliveFuture ());
971+ }
972+
973+ @ Test
974+ public void testKeepAliveRunnableHandlesSynchronousException () {
975+ DatabaseClient client = mock (DatabaseClient .class );
976+ when (client .transactionManager ())
977+ .thenAnswer (
978+ invocation -> {
979+ TransactionContext txContext = mock (TransactionContext .class );
980+ when (txContext .executeQuery (any (Statement .class )))
981+ .thenThrow (new RuntimeException ("Simulated synchronous execution error" ));
982+ return new SimpleTransactionManager (txContext , CommitBehavior .SUCCEED );
983+ });
984+
985+ ReadWriteTransaction transaction =
986+ ReadWriteTransaction .newBuilder ()
987+ .setDatabaseClient (client )
988+ .setKeepTransactionAlive (true )
989+ .setRetryAbortsInternally (false )
990+ .setIsolationLevel (IsolationLevel .ISOLATION_LEVEL_UNSPECIFIED )
991+ .setSavepointSupport (SavepointSupport .FAIL_AFTER_ROLLBACK )
992+ .setTransactionRetryListeners (Collections .emptyList ())
993+ .withStatementExecutor (new StatementExecutor ())
994+ .setSpan (Span .getInvalid ())
995+ .build ();
996+
997+ ReadWriteTransaction .KeepAliveRunnable runnable =
998+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
999+
1000+ runnable .run ();
1001+
1002+ assertFalse (transaction .abortedLock .isLocked ());
1003+ }
1004+
8601005 private static StatusRuntimeException createAbortedExceptionWithMinimalRetry () {
8611006 Metadata .Key <RetryInfo > key = ProtoUtils .keyForProto (RetryInfo .getDefaultInstance ());
8621007 Metadata trailers = new Metadata ();
0 commit comments