Fix watch leak when watch is set after subscription is disposed

This commit is contained in:
ameya-signal authored and GitHub committed 2026-07-21 10:28:33 -04:00
1 parent c082b76729
commit 2cbc33f136
2 files changed
+63

No files matched your search

@@ -377,6 +377,10 @@ class FoundationDbMessagePublisher {
private synchronized void setWatch(final Transaction transaction) {
assert messagesAvailableWatchKey != null;
if (state == State.TERMINATED || state == State.ERROR) {
// If we're already in a terminal state, we don't want to set a watch that would leak
return;
}
// Set a watch that will be triggered when new messages arrive. When it is triggered, we attempt to fetch messages
// again (if there is demand). When we run out of messages, this method will be called again, setting another watch,
// and so on, thus achieving a "watch for new messages -> read -> publish" loop.
@@ -9,6 +9,7 @@ import static org.mockito.ArgumentMatchers.anyInt;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.verifyNoInteractions;
import static org.mockito.Mockito.when;
import com.apple.foundationdb.Database;
@@ -45,6 +46,7 @@ import org.whispersystems.textsecuregcm.entities.MessageProtos;
import org.whispersystems.textsecuregcm.identity.AciServiceIdentifier;
import org.whispersystems.textsecuregcm.storage.Device;
import org.whispersystems.textsecuregcm.storage.MessageStreamEntry;
import reactor.core.Disposable;
import reactor.test.StepVerifier;
/// NOTE: most of the happy-path test cases are already covered in {@link FoundationDbMessageStoreTest}, this test
@@ -478,6 +480,63 @@ class FoundationDbMessagePublisherTest {
}
@Test
@SuppressWarnings({"unchecked"})
void watchSetAfterTerminationIsCancelled() {
final Transaction transaction = mock(Transaction.class);
final CompletableFuture<List<KeyValue>> batchFuture = new CompletableFuture<>();
final AsyncIterable<KeyValue> batch = mock(AsyncIterable.class);
when(batch.asList()).thenReturn(batchFuture);
when(transaction.getRange(any(KeySelector.class), any(KeySelector.class), anyInt(), anyBoolean(),
any(StreamingMode.class)))
.thenReturn(batch);
final CompletableFuture<Void> watchFuture = mock(CompletableFuture.class);
when(transaction.watch(MESSAGES_AVAILABLE_WATCH_KEY)).thenReturn(watchFuture);
when(transaction.get(FoundationDbMessageStore.getPresenceKey(SERVICE_IDENTIFIER, Device.PRIMARY_ID)))
.thenReturn(CompletableFuture.completedFuture(null));
when(database.runAsync(any(Function.class))).thenAnswer(
(Answer<CompletableFuture<List<? extends MessageStreamEntry>>>) invocationOnMock -> {
final Function<Transaction, CompletableFuture<List<? extends MessageStreamEntry>>> f =
invocationOnMock.getArgument(0);
return f.apply(transaction);
});
final FoundationDbMessagePublisher infinitePublisher = new FoundationDbMessagePublisher(
database,
CLOCK,
KeySelector.firstGreaterOrEqual(SUBSPACE_RANGE.begin),
KeySelector.firstGreaterOrEqual(SUBSPACE_RANGE.end),
1,
FoundationDbMessageStore.getPresenceKey(SERVICE_IDENTIFIER, Device.PRIMARY_ID),
presenceRenewalExecutorService,
MESSAGES_AVAILABLE_WATCH_KEY,
(_, newState) -> stateTransitions.add(newState)
);
// Subscribe requests demand synchronously, and the publisher begins fetching messages; however we control the fetch
// result future, so the state here should be FETCHING_MESSAGES
final Disposable subscription = infinitePublisher.getMessages().subscribe(_ -> {}, _ -> {});
verify(transaction, times(1)).getRange(any(KeySelector.class), any(KeySelector.class), anyInt(), anyBoolean(),
any(StreamingMode.class));
// Dispose the subscription while the fetch is still in flight
subscription.dispose();
// Complete the fetch result now, which should normally set the watch
batchFuture.complete(List.of());
// Verify that the watch was never set
verify(transaction, times(0)).watch(MESSAGES_AVAILABLE_WATCH_KEY);
verifyNoInteractions(watchFuture);
assertEquals(List.of(
FoundationDbMessagePublisher.State.FETCHING_MESSAGES,
FoundationDbMessagePublisher.State.TERMINATED
), stateTransitions);
}
private FoundationDbMessageStreamEntry.Message getExpectedMessageStreamEntry(final KeyValue keyValue) {
try {
return new FoundationDbMessageStreamEntry.Message(FoundationDbMessageStore.getVersionstamp(keyValue.getKey()),