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,172 @@ 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 = waitForKeepAliveFuture (transaction );
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 = waitForKeepAliveFuture (transaction );
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 (waitForKeepAliveFuture (transaction ));
969+
970+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
971+
972+ assertNull (transaction .getKeepAliveFuture ());
973+ }
974+
975+ private static ScheduledFuture <?> waitForKeepAliveFuture (ReadWriteTransaction transaction ) {
976+ long deadline = System .currentTimeMillis () + 5000 ;
977+ while (System .currentTimeMillis () < deadline ) {
978+ ScheduledFuture <?> future = transaction .getKeepAliveFuture ();
979+ if (future != null ) {
980+ return future ;
981+ }
982+ try {
983+ Thread .sleep (1 );
984+ } catch (InterruptedException e ) {
985+ Thread .currentThread ().interrupt ();
986+ break ;
987+ }
988+ }
989+ fail ("Keep-alive future was not populated within 5 seconds" );
990+ return null ;
991+ }
992+
993+ @ Test
994+ public void testKeepAliveRunnableHandlesSynchronousException () {
995+ DatabaseClient client = mock (DatabaseClient .class );
996+ when (client .transactionManager ())
997+ .thenAnswer (
998+ invocation -> {
999+ TransactionContext txContext = mock (TransactionContext .class );
1000+ when (txContext .executeQuery (any (Statement .class )))
1001+ .thenThrow (new RuntimeException ("Simulated synchronous execution error" ));
1002+ return new SimpleTransactionManager (txContext , CommitBehavior .SUCCEED );
1003+ });
1004+
1005+ ReadWriteTransaction transaction =
1006+ ReadWriteTransaction .newBuilder ()
1007+ .setDatabaseClient (client )
1008+ .setKeepTransactionAlive (true )
1009+ .setRetryAbortsInternally (false )
1010+ .setIsolationLevel (IsolationLevel .ISOLATION_LEVEL_UNSPECIFIED )
1011+ .setSavepointSupport (SavepointSupport .FAIL_AFTER_ROLLBACK )
1012+ .setTransactionRetryListeners (Collections .emptyList ())
1013+ .withStatementExecutor (new StatementExecutor ())
1014+ .setSpan (Span .getInvalid ())
1015+ .build ();
1016+
1017+ ReadWriteTransaction .KeepAliveRunnable runnable =
1018+ new ReadWriteTransaction .KeepAliveRunnable (transaction );
1019+
1020+ runnable .run ();
1021+
1022+ assertFalse (transaction .abortedLock .isLocked ());
1023+ }
1024+
1025+ @ Test
1026+ public void testKeepAliveNotScheduledIfTransactionClosed () {
1027+ ParsedStatement parsedStatement = mock (ParsedStatement .class );
1028+ when (parsedStatement .getType ()).thenReturn (StatementType .UPDATE );
1029+ when (parsedStatement .isUpdate ()).thenReturn (true );
1030+ Statement statement = Statement .of ("UPDATE FOO SET BAR=1 WHERE ID=2" );
1031+ when (parsedStatement .getStatement ()).thenReturn (statement );
1032+
1033+ ReadWriteTransaction transaction = createSubject (/* keepTransactionAlive= */ true );
1034+ get (transaction .executeUpdateAsync (CallType .SYNC , parsedStatement ));
1035+ get (transaction .commitAsync (CallType .SYNC , NoopEndTransactionCallback .INSTANCE ));
1036+
1037+ transaction .maybeScheduleKeepAlivePing ();
1038+
1039+ assertNull (transaction .getKeepAliveFuture ());
1040+ }
1041+
8601042 private static StatusRuntimeException createAbortedExceptionWithMinimalRetry () {
8611043 Metadata .Key <RetryInfo > key = ProtoUtils .keyForProto (RetryInfo .getDefaultInstance ());
8621044 Metadata trailers = new Metadata ();
0 commit comments