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 ;
32+ import static org .junit .Assert .assertTrue ;
2933import static org .junit .Assert .fail ;
3034import static org .mockito .Mockito .any ;
3135import static org .mockito .Mockito .doThrow ;
6771import java .math .BigDecimal ;
6872import java .util .Arrays ;
6973import java .util .Collections ;
74+ import java .util .concurrent .CountDownLatch ;
75+ import java .util .concurrent .ScheduledFuture ;
7076import org .junit .Test ;
7177import org .junit .runner .RunWith ;
7278import org .junit .runners .JUnit4 ;
@@ -158,12 +164,21 @@ private ReadWriteTransaction createSubject() {
158164 return createSubject (CommitBehavior .SUCCEED , false );
159165 }
160166
167+ private ReadWriteTransaction createSubject (boolean keepTransactionAlive ) {
168+ return createSubject (CommitBehavior .SUCCEED , false , keepTransactionAlive );
169+ }
170+
161171 private ReadWriteTransaction createSubject (CommitBehavior commitBehavior ) {
162- return createSubject (commitBehavior , false );
172+ return createSubject (commitBehavior , false , false );
163173 }
164174
165175 private ReadWriteTransaction createSubject (
166176 final CommitBehavior commitBehavior , boolean withRetry ) {
177+ return createSubject (commitBehavior , withRetry , false );
178+ }
179+
180+ private ReadWriteTransaction createSubject (
181+ final CommitBehavior commitBehavior , boolean withRetry , boolean keepTransactionAlive ) {
167182 DatabaseClient client = mock (DatabaseClient .class );
168183 when (client .transactionManager ())
169184 .thenAnswer (
@@ -179,6 +194,7 @@ private ReadWriteTransaction createSubject(
179194 });
180195 return ReadWriteTransaction .newBuilder ()
181196 .setDatabaseClient (client )
197+ .setKeepTransactionAlive (keepTransactionAlive )
182198 .setRetryAbortsInternally (withRetry )
183199 .setIsolationLevel (IsolationLevel .ISOLATION_LEVEL_UNSPECIFIED )
184200 .setSavepointSupport (SavepointSupport .FAIL_AFTER_ROLLBACK )
@@ -857,6 +873,137 @@ public void testGetCommitResponseAfterCommit() {
857873 assertNotNull (transaction .getCommitResponseOrNull ());
858874 }
859875
876+ @ Test
877+ public void testKeepAliveTaskRemovedFromQueueOnCancel () {
878+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
879+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
880+ when (parsedStatement .isUpdate ()).thenReturn (true );
881+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
882+ when (parsedStatement .getStatement ()).thenReturn (statement );
883+
884+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
885+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
886+
887+ ScheduledFuture <?> future = transaction .getKeepAliveFuture ();
888+ assertNotNull (future );
889+ assertTrue (ReadWriteTransaction .getKeepAliveService ().getQueue ().contains (future ));
890+
891+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
892+ assertFalse (ReadWriteTransaction .getKeepAliveService ().getQueue ().contains (future ));
893+ }
894+
895+ @ Test
896+ public void testKeepAliveWeakReference () {
897+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
898+ ReadWriteTransaction .KeepAliveRunnable runnable =
899+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
900+
901+ assertNotNull (runnable .transactionRef );
902+ assertSame (transaction , runnable .transactionRef .get ());
903+ }
904+
905+ @ Test
906+ public void testKeepAliveRescheduledWhenLockBusy () {
907+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
908+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
909+ when (parsedStatement .isUpdate ()).thenReturn (true );
910+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
911+ when (parsedStatement .getStatement ()).thenReturn (statement );
912+
913+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
914+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
915+
916+ ScheduledFuture <?> future1 = transaction .getKeepAliveFuture ();
917+ assertNotNull (future1 );
918+
919+ CountDownLatch latch = new CountDownLatch (1 );
920+ CountDownLatch lockAcquired = new CountDownLatch (1 );
921+ Thread lockHoldingThread =
922+ new Thread (
923+ () -> {
924+ transaction .abortedLock .lock ();
925+ try {
926+ lockAcquired .countDown ();
927+ latch .await ();
928+ } catch (InterruptedException e ) {
929+ Thread .currentThread ().interrupt ();
930+ } finally {
931+ transaction .abortedLock .unlock ();
932+ }
933+ });
934+ lockHoldingThread .start ();
935+ try {
936+ lockAcquired .await ();
937+ ReadWriteTransaction .KeepAliveRunnable runnable =
938+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
939+ runnable .run ();
940+ } catch (InterruptedException e ) {
941+ Thread .currentThread ().interrupt ();
942+ fail ("Test interrupted" );
943+ } finally {
944+ latch .countDown ();
945+ try {
946+ lockHoldingThread .join ();
947+ } catch (InterruptedException e ) {
948+ Thread .currentThread ().interrupt ();
949+ }
950+ }
951+
952+ ScheduledFuture <?> future2 = transaction .getKeepAliveFuture ();
953+ assertNotNull (future2 );
954+ assertNotSame (future1 , future2 );
955+ }
956+
957+ @ Test
958+ public void testKeepAliveFutureNullifiedOnCancel () {
959+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
960+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
961+ when (parsedStatement .isUpdate ()).thenReturn (true );
962+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
963+ when (parsedStatement .getStatement ()).thenReturn (statement );
964+
965+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
966+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
967+
968+ assertNotNull (transaction .getKeepAliveFuture ());
969+
970+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
971+
972+ assertNull (transaction .getKeepAliveFuture ());
973+ }
974+
975+ @ Test
976+ public void testKeepAliveRunnableHandlesSynchronousException () {
977+ DatabaseClient client = mock (DatabaseClient .class );
978+ when (client .transactionManager ())
979+ .thenAnswer (
980+ invocation -> {
981+ TransactionContext txContext = mock (TransactionContext .class );
982+ when (txContext .executeQuery (any (Statement .class )))
983+ .thenThrow (new RuntimeException ("Simulated synchronous execution error" ));
984+ return new SimpleTransactionManager (txContext , CommitBehavior .SUCCEED );
985+ });
986+
987+ ReadWriteTransaction transaction =
988+ ReadWriteTransaction .newBuilder ()
989+ .setDatabaseClient (client )
990+ .setKeepTransactionAlive (true )
991+ .setRetryAbortsInternally (false )
992+ .setIsolationLevel (IsolationLevel .ISOLATION_LEVEL_UNSPECIFIED )
993+ .setSavepointSupport (SavepointSupport .FAIL_AFTER_ROLLBACK )
994+ .setTransactionRetryListeners (Collections .emptyList ())
995+ .withStatementExecutor (new StatementExecutor ())
996+ .setSpan (Span .getInvalid ())
997+ .build ();
998+
999+ ReadWriteTransaction .KeepAliveRunnable runnable =
1000+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
1001+
1002+ runnable .run ();
1003+
1004+ assertFalse (transaction .abortedLock .isLocked ());
1005+ }
1006+
8601007 private static StatusRuntimeException createAbortedExceptionWithMinimalRetry () {
8611008 Metadata .Key <RetryInfo > key = ProtoUtils .keyForProto (RetryInfo .getDefaultInstance ());
8621009 Metadata trailers = new Metadata ();
0 commit comments