From ef6529fe09810769a5b581575d6e0cdb46c56334 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Knut=20Olav=20L=C3=B8ite?= Date: Mon, 27 Jul 2026 11:07:27 +0200 Subject: [PATCH] fix(spanner): prevent memory leak and thread blocking in transaction keep-alive - Enable `setRemoveOnCancelPolicy(true)` on `KEEP_ALIVE_SERVICE` so canceled tasks are immediately purged from `DelayedWorkQueue`. - Use a `WeakReference` in `KeepAliveRunnable` to prevent scheduled tasks from retaining strong references to transaction instances. - Use `abortedLock.tryLock()` in `KeepAliveRunnable` so the shared executor thread does not block when a transaction is active or retrying. - Remove duplicate `maybeScheduleKeepAlivePing` listener registration on keep-alive query completion. - Add unit tests in `ReadWriteTransactionTest` verifying task removal on cancel, weak reference retention, non-blocking lock handling, and single ping scheduling on completion. --- .../connection/ReadWriteTransaction.java | 75 +++++++-- .../connection/ReadWriteTransactionTest.java | 148 +++++++++++++++++- 2 files changed, 208 insertions(+), 15 deletions(-) diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/connection/ReadWriteTransaction.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/connection/ReadWriteTransaction.java index ccb592e3f843..3a355cc62c17 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/connection/ReadWriteTransaction.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/connection/ReadWriteTransaction.java @@ -65,15 +65,15 @@ import io.grpc.Deadline; import io.opentelemetry.api.common.AttributeKey; import io.opentelemetry.context.Scope; +import java.lang.ref.WeakReference; import java.time.Duration; import java.util.ArrayList; import java.util.LinkedList; import java.util.List; import java.util.Objects; import java.util.concurrent.Callable; -import java.util.concurrent.Executors; -import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.ScheduledThreadPoolExecutor; import java.util.concurrent.ThreadFactory; import java.util.concurrent.ThreadLocalRandom; import java.util.concurrent.TimeUnit; @@ -100,8 +100,20 @@ class ReadWriteTransaction extends AbstractMultiUseTransaction { private static final ThreadFactory KEEP_ALIVE_THREAD_FACTORY = ThreadFactoryUtil.createVirtualOrPlatformDaemonThreadFactory( "read-write-transaction-keep-alive", true); - private static final ScheduledExecutorService KEEP_ALIVE_SERVICE = - Executors.newSingleThreadScheduledExecutor(KEEP_ALIVE_THREAD_FACTORY); + private static final ScheduledThreadPoolExecutor KEEP_ALIVE_SERVICE = createKeepAliveService(); + + private static ScheduledThreadPoolExecutor createKeepAliveService() { + ScheduledThreadPoolExecutor executor = + new ScheduledThreadPoolExecutor(1, KEEP_ALIVE_THREAD_FACTORY); + executor.setRemoveOnCancelPolicy(true); + return executor; + } + + @VisibleForTesting + static ScheduledThreadPoolExecutor getKeepAliveService() { + return KEEP_ALIVE_SERVICE; + } + private static final ParsedStatement SELECT1_STATEMENT = AbstractStatementParser.getInstance(Dialect.GOOGLE_STANDARD_SQL) .parse(Statement.of("SELECT 1")); @@ -146,7 +158,7 @@ class ReadWriteTransaction extends AbstractMultiUseTransaction { private Savepoint autoSavepoint; private final int maxInternalRetries; - private final ReentrantLock abortedLock = new ReentrantLock(); + final ReentrantLock abortedLock = new ReentrantLock(); private final long transactionId; private final DatabaseClient dbClient; private final TransactionOption[] transactionOptions; @@ -475,7 +487,7 @@ private void maybeScheduleKeepAlivePing() { if (keepAliveFuture == null || keepAliveFuture.isDone()) { keepAliveFuture = KEEP_ALIVE_SERVICE.schedule( - new KeepAliveRunnable(), + new KeepAliveRunnable(this), keepAliveIntervalMillis > 0 ? keepAliveIntervalMillis : DEFAULT_KEEP_ALIVE_INTERVAL_MILLIS, @@ -487,12 +499,18 @@ private void maybeScheduleKeepAlivePing() { } } + @VisibleForTesting + ScheduledFuture getKeepAliveFuture() { + return keepAliveFuture; + } + private void cancelScheduledKeepAlivePing() { if (keepAliveLock != null) { keepAliveLock.lock(); try { if (keepAliveFuture != null) { keepAliveFuture.cancel(false); + keepAliveFuture = null; } } finally { keepAliveLock.unlock(); @@ -500,14 +518,35 @@ private void cancelScheduledKeepAlivePing() { } } - private class KeepAliveRunnable implements Runnable { + private void rescheduleKeepAlivePing() { + if (keepAliveLock != null) { + keepAliveLock.lock(); + try { + keepAliveFuture = null; + maybeScheduleKeepAlivePing(); + } finally { + keepAliveLock.unlock(); + } + } + } + + static class KeepAliveRunnable implements Runnable { + final WeakReference transactionRef; + + KeepAliveRunnable(ReadWriteTransaction transaction) { + this.transactionRef = new WeakReference<>(transaction); + } + @Override public void run() { - if (shouldPing()) { - // Do a shoot-and-forget ping and schedule a new ping over 8 seconds after this ping has - // finished. - ApiFuture future = - executeQueryAsync( + ReadWriteTransaction transaction = transactionRef.get(); + if (transaction != null && transaction.shouldPing()) { + if (transaction.abortedLock.tryLock()) { + try { + // Do a shoot-and-forget ping. + // Note: executeQueryAsync automatically adds StatementResultCallback, + // which calls maybeScheduleKeepAlivePing() upon completion. + transaction.executeQueryAsync( CallType.SYNC, SELECT1_STATEMENT, AnalyzeMode.NONE, @@ -515,8 +554,16 @@ public void run() { System.getProperty( "spanner.connection.keep_alive_query_tag", "connection.transaction-keep-alive"))); - future.addListener( - ReadWriteTransaction.this::maybeScheduleKeepAlivePing, MoreExecutors.directExecutor()); + } catch (Throwable t) { + transaction.maybeScheduleKeepAlivePing(); + } finally { + transaction.abortedLock.unlock(); + } + } else { + // Transaction is currently busy (executing a statement or retrying). + // Reschedule keep-alive ping for later since it is active. + transaction.rescheduleKeepAlivePing(); + } } } } diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/connection/ReadWriteTransactionTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/connection/ReadWriteTransactionTest.java index 7d0fa94c9b0b..2edcfbf40f8d 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/connection/ReadWriteTransactionTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/connection/ReadWriteTransactionTest.java @@ -24,8 +24,11 @@ import static org.hamcrest.CoreMatchers.nullValue; import static org.hamcrest.MatcherAssert.assertThat; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertNotSame; import static org.junit.Assert.assertNull; +import static org.junit.Assert.assertSame; import static org.junit.Assert.fail; import static org.mockito.Mockito.any; import static org.mockito.Mockito.doThrow; @@ -67,6 +70,8 @@ import java.math.BigDecimal; import java.util.Arrays; import java.util.Collections; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ScheduledFuture; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.JUnit4; @@ -158,12 +163,21 @@ private ReadWriteTransaction createSubject() { return createSubject(CommitBehavior.SUCCEED, false); } + private ReadWriteTransaction createSubject(boolean keepTransactionAlive) { + return createSubject(CommitBehavior.SUCCEED, false, keepTransactionAlive); + } + private ReadWriteTransaction createSubject(CommitBehavior commitBehavior) { - return createSubject(commitBehavior, false); + return createSubject(commitBehavior, false, false); } private ReadWriteTransaction createSubject( final CommitBehavior commitBehavior, boolean withRetry) { + return createSubject(commitBehavior, withRetry, false); + } + + private ReadWriteTransaction createSubject( + final CommitBehavior commitBehavior, boolean withRetry, boolean keepTransactionAlive) { DatabaseClient client = mock(DatabaseClient.class); when(client.transactionManager()) .thenAnswer( @@ -179,6 +193,7 @@ private ReadWriteTransaction createSubject( }); return ReadWriteTransaction.newBuilder() .setDatabaseClient(client) + .setKeepTransactionAlive(keepTransactionAlive) .setRetryAbortsInternally(withRetry) .setIsolationLevel(IsolationLevel.ISOLATION_LEVEL_UNSPECIFIED) .setSavepointSupport(SavepointSupport.FAIL_AFTER_ROLLBACK) @@ -857,6 +872,137 @@ public void testGetCommitResponseAfterCommit() { assertNotNull(transaction.getCommitResponseOrNull()); } + @Test + public void testKeepAliveTaskRemovedFromQueueOnCancel() { + ParsedStatement parsedStatement = mock(ParsedStatement.class); + when(parsedStatement.getType()).thenReturn(StatementType.UPDATE); + when(parsedStatement.isUpdate()).thenReturn(true); + Statement statement = Statement.of("UPDATE FOO SET BAR=1 WHERE ID=2"); + when(parsedStatement.getStatement()).thenReturn(statement); + + int initialQueueSize = ReadWriteTransaction.getKeepAliveService().getQueue().size(); + ReadWriteTransaction transaction = createSubject(/* keepTransactionAlive= */ true); + get(transaction.executeUpdateAsync(CallType.SYNC, parsedStatement)); + + assertEquals( + initialQueueSize + 1, ReadWriteTransaction.getKeepAliveService().getQueue().size()); + + get(transaction.commitAsync(CallType.SYNC, NoopEndTransactionCallback.INSTANCE)); + assertEquals(initialQueueSize, ReadWriteTransaction.getKeepAliveService().getQueue().size()); + } + + @Test + public void testKeepAliveWeakReference() { + ReadWriteTransaction transaction = createSubject(/* keepTransactionAlive= */ true); + ReadWriteTransaction.KeepAliveRunnable runnable = + new ReadWriteTransaction.KeepAliveRunnable(transaction); + + assertNotNull(runnable.transactionRef); + assertSame(transaction, runnable.transactionRef.get()); + } + + @Test + public void testKeepAliveRescheduledWhenLockBusy() { + ParsedStatement parsedStatement = mock(ParsedStatement.class); + when(parsedStatement.getType()).thenReturn(StatementType.UPDATE); + when(parsedStatement.isUpdate()).thenReturn(true); + Statement statement = Statement.of("UPDATE FOO SET BAR=1 WHERE ID=2"); + when(parsedStatement.getStatement()).thenReturn(statement); + + ReadWriteTransaction transaction = createSubject(/* keepTransactionAlive= */ true); + get(transaction.executeUpdateAsync(CallType.SYNC, parsedStatement)); + + ScheduledFuture future1 = transaction.getKeepAliveFuture(); + assertNotNull(future1); + + CountDownLatch latch = new CountDownLatch(1); + CountDownLatch lockAcquired = new CountDownLatch(1); + Thread lockHoldingThread = + new Thread( + () -> { + transaction.abortedLock.lock(); + try { + lockAcquired.countDown(); + latch.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } finally { + transaction.abortedLock.unlock(); + } + }); + lockHoldingThread.start(); + try { + lockAcquired.await(); + ReadWriteTransaction.KeepAliveRunnable runnable = + new ReadWriteTransaction.KeepAliveRunnable(transaction); + runnable.run(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + fail("Test interrupted"); + } finally { + latch.countDown(); + try { + lockHoldingThread.join(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + + ScheduledFuture future2 = transaction.getKeepAliveFuture(); + assertNotNull(future2); + assertNotSame(future1, future2); + } + + @Test + public void testKeepAliveFutureNullifiedOnCancel() { + ParsedStatement parsedStatement = mock(ParsedStatement.class); + when(parsedStatement.getType()).thenReturn(StatementType.UPDATE); + when(parsedStatement.isUpdate()).thenReturn(true); + Statement statement = Statement.of("UPDATE FOO SET BAR=1 WHERE ID=2"); + when(parsedStatement.getStatement()).thenReturn(statement); + + ReadWriteTransaction transaction = createSubject(/* keepTransactionAlive= */ true); + get(transaction.executeUpdateAsync(CallType.SYNC, parsedStatement)); + + assertNotNull(transaction.getKeepAliveFuture()); + + get(transaction.commitAsync(CallType.SYNC, NoopEndTransactionCallback.INSTANCE)); + + assertNull(transaction.getKeepAliveFuture()); + } + + @Test + public void testKeepAliveRunnableHandlesSynchronousException() { + DatabaseClient client = mock(DatabaseClient.class); + when(client.transactionManager()) + .thenAnswer( + invocation -> { + TransactionContext txContext = mock(TransactionContext.class); + when(txContext.executeQuery(any(Statement.class))) + .thenThrow(new RuntimeException("Simulated synchronous execution error")); + return new SimpleTransactionManager(txContext, CommitBehavior.SUCCEED); + }); + + ReadWriteTransaction transaction = + ReadWriteTransaction.newBuilder() + .setDatabaseClient(client) + .setKeepTransactionAlive(true) + .setRetryAbortsInternally(false) + .setIsolationLevel(IsolationLevel.ISOLATION_LEVEL_UNSPECIFIED) + .setSavepointSupport(SavepointSupport.FAIL_AFTER_ROLLBACK) + .setTransactionRetryListeners(Collections.emptyList()) + .withStatementExecutor(new StatementExecutor()) + .setSpan(Span.getInvalid()) + .build(); + + ReadWriteTransaction.KeepAliveRunnable runnable = + new ReadWriteTransaction.KeepAliveRunnable(transaction); + + runnable.run(); + + assertFalse(transaction.abortedLock.isLocked()); + } + private static StatusRuntimeException createAbortedExceptionWithMinimalRetry() { Metadata.Key key = ProtoUtils.keyForProto(RetryInfo.getDefaultInstance()); Metadata trailers = new Metadata();