From 740116e1fd2f5be4b9ff32e29934e829405dc6c7 Mon Sep 17 00:00:00 2001 From: Jon Chambers Date: Tue, 22 Sep 2026 17:21:37 -0400 Subject: [PATCH] Don't insert shared MRM payloads for unresolved recipients --- .../textsecuregcm/storage/MessagesCache.java | 6 ++- ...edMultiRecipientPayloadAndViewsScript.java | 24 +++++++-- .../storage/MessagesManager.java | 35 +++++++----- ...ltiRecipientPayloadAndViewsScriptTest.java | 54 +++++++++++++++++-- ...oveRecipientViewFromMrmDataScriptTest.java | 12 ++++- .../storage/MessagesCacheTest.java | 6 +-- .../storage/MessagesManagerTest.java | 2 +- 7 files changed, 108 insertions(+), 31 deletions(-) diff --git a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCache.java b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCache.java index a306538bc..481fbd418 100644 --- a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCache.java +++ b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCache.java @@ -35,6 +35,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.Set; import java.util.UUID; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutorService; @@ -236,13 +237,14 @@ public class MessagesCache { } public CompletableFuture insertSharedMultiRecipientMessagePayload( - final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage) { + final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage, + final Set resolvedRecipients) { final Timer.Sample sample = Timer.start(); final byte[] sharedMrmKey = getSharedMrmKey(UUID.randomUUID()); - return insertMrmScript.executeAsync(sharedMrmKey, sealedSenderMultiRecipientMessage) + return insertMrmScript.executeAsync(sharedMrmKey, sealedSenderMultiRecipientMessage, resolvedRecipients) .thenApply(_ -> sharedMrmKey) .toCompletableFuture() .whenComplete((_, _) -> sample.stop(insertSharedMrmPayloadTimer)); diff --git a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript.java b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript.java index 8b435b00a..a56755ce3 100644 --- a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript.java +++ b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript.java @@ -9,6 +9,8 @@ import io.lettuce.core.ScriptOutputType; import java.io.IOException; import java.util.ArrayList; import java.util.List; +import java.util.Map; +import java.util.Set; import java.util.concurrent.CompletionStage; import java.util.concurrent.ScheduledExecutorService; import org.signal.libsignal.protocol.SealedSenderMultiRecipientMessage; @@ -39,18 +41,30 @@ class MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript { this.retryExecutor = retryExecutor; } - CompletionStage executeAsync(final byte[] sharedMrmKey, final SealedSenderMultiRecipientMessage message) { + CompletionStage executeAsync(final byte[] sharedMrmKey, + final SealedSenderMultiRecipientMessage message, + final Set resolvedRecipients) { + final List keys = List.of( sharedMrmKey // sharedMrmKey ); - // Pre-allocate capacity for the most fields we expect -- 6 devices per recipient, plus the data field. - final List args = new ArrayList<>(message.getRecipients().size() * 6 + 1); + final int deviceCount = message.getRecipients().values().stream() + .filter(resolvedRecipients::contains) + .mapToInt(recipient -> recipient.getDevices().length) + .sum(); + + // Pre-allocate capacity for the most fields we expect -- the shared data field plus one for each device + final List args = new ArrayList<>(1 + deviceCount); args.add(message.serialized()); message.getRecipients().forEach((serviceId, recipient) -> { - for (byte device : recipient.getDevices()) { - args.add(MessagesCache.getSharedMrmViewKey(serviceId, device)); + if (!resolvedRecipients.contains(recipient)) { + return; + } + + for (final byte deviceId : recipient.getDevices()) { + args.add(MessagesCache.getSharedMrmViewKey(serviceId, deviceId)); args.add(message.serializedRecipientView(recipient)); } }); diff --git a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesManager.java b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesManager.java index c42a46423..a6ef3ece8 100644 --- a/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesManager.java +++ b/service/src/main/java/org/whispersystems/textsecuregcm/storage/MessagesManager.java @@ -17,6 +17,7 @@ import java.util.Collections; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.Set; import java.util.UUID; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; @@ -229,15 +230,17 @@ public class MessagesManager { final long serverTimestamp = clock.millis(); - return insertSharedMultiRecipientMessagePayload(multiRecipientMessage) - .thenCompose(sharedMrmKey -> { + return insertSharedMultiRecipientMessagePayload(multiRecipientMessage, resolvedRecipients.keySet()) + .thenCompose(maybeSharedMrmKey -> { final Envelope.Builder envelopeBuilder = Envelope.newBuilder() .setType(Envelope.Type.UNIDENTIFIED_SENDER) .setClientTimestamp(clientTimestamp == 0 ? serverTimestamp : clientTimestamp) .setServerTimestamp(serverTimestamp) .setEphemeral(isEphemeral) - .setUrgent(isUrgent) - .setSharedMrmKey(ByteString.copyFrom(sharedMrmKey)); + .setUrgent(isUrgent); + + maybeSharedMrmKey + .ifPresent(sharedMrmKey -> envelopeBuilder.setSharedMrmKey(ByteString.copyFrom(sharedMrmKey))); if (isStory) { // Avoid sending this field if it's false. @@ -407,15 +410,21 @@ public class MessagesManager { .toFuture(); } - /** - * Inserts the shared multi-recipient message payload to storage. - * - * @return a key where the shared data is stored - * @see MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript - */ - private CompletableFuture insertSharedMultiRecipientMessagePayload( - final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage) { - return messagesCache.insertSharedMultiRecipientMessagePayload(sealedSenderMultiRecipientMessage); + /// Inserts the shared multi-recipient message payload to storage. + /// + /// @return a future that yields a key where the shared data is stored or empty if the resolved recipient set is empty + /// + /// @see MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript + private CompletableFuture> insertSharedMultiRecipientMessagePayload( + final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage, + final Set resolvedRecipients) { + + if (resolvedRecipients.isEmpty()) { + return CompletableFuture.completedFuture(Optional.empty()); + } + + return messagesCache.insertSharedMultiRecipientMessagePayload(sealedSenderMultiRecipientMessage, resolvedRecipients) + .thenApply(Optional::of); } /// Record versionstamps for the current time in the FoundationDB database(s). diff --git a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest.java b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest.java index e6eea0a21..be600ee23 100644 --- a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest.java +++ b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest.java @@ -12,10 +12,13 @@ import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.mock; import io.lettuce.core.RedisException; +import java.io.IOException; import java.util.ArrayList; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Set; import java.util.UUID; import java.util.concurrent.CompletionException; import java.util.concurrent.ScheduledExecutorService; @@ -26,6 +29,7 @@ import org.junit.jupiter.api.extension.RegisterExtension; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; +import org.signal.libsignal.protocol.SealedSenderMultiRecipientMessage; import org.whispersystems.textsecuregcm.identity.AciServiceIdentifier; import org.whispersystems.textsecuregcm.identity.ServiceIdentifier; import org.whispersystems.textsecuregcm.redis.RedisClusterExtension; @@ -43,9 +47,12 @@ class MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest { final MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript insertMrmScript = new MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript( REDIS_CLUSTER_EXTENSION.getRedisCluster(), mock(ScheduledExecutorService.class)); + final SealedSenderMultiRecipientMessage multiRecipientMessage = + MessagesCacheTest.generateRandomMrmMessage(destinations); + final byte[] sharedMrmKey = MessagesCache.getSharedMrmKey(UUID.randomUUID()); - insertMrmScript.executeAsync(sharedMrmKey, - MessagesCacheTest.generateRandomMrmMessage(destinations)).toCompletableFuture().join(); + insertMrmScript.executeAsync(sharedMrmKey, multiRecipientMessage, new HashSet<>(multiRecipientMessage.getRecipients().values())) + .toCompletableFuture().join(); final int totalDevices = destinations.values().stream().mapToInt(List::size).sum(); final long hashFieldCount = REDIS_CLUSTER_EXTENSION.getRedisCluster() @@ -81,21 +88,58 @@ class MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScriptTest { return testCases; } + @Test + void testInsertUnresolvedRecipient() throws IOException { + final MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript insertMrmScript = new MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript( + REDIS_CLUSTER_EXTENSION.getRedisCluster(), mock(ScheduledExecutorService.class)); + + final AciServiceIdentifier resolvedServiceIdentifier = new AciServiceIdentifier(UUID.randomUUID()); + final AciServiceIdentifier unresolvedServiceIdentifier = new AciServiceIdentifier(UUID.randomUUID()); + + final Map> destinations = Map.of( + resolvedServiceIdentifier, List.of(Device.PRIMARY_ID), + unresolvedServiceIdentifier, List.of(Device.PRIMARY_ID)); + + final SealedSenderMultiRecipientMessage multiRecipientMessage = + MessagesCacheTest.generateRandomMrmMessage(destinations); + + final Set resolvedRecipients = + multiRecipientMessage.getRecipients().entrySet().stream() + .filter(entry -> entry.getKey().getRawUUID().equals(resolvedServiceIdentifier.uuid())) + .map(Map.Entry::getValue) + .collect(Collectors.toSet()); + + final byte[] sharedMrmKey = MessagesCache.getSharedMrmKey(UUID.randomUUID()); + insertMrmScript.executeAsync(sharedMrmKey, multiRecipientMessage, resolvedRecipients) + .toCompletableFuture().join(); + + final long hashFieldCount = REDIS_CLUSTER_EXTENSION.getRedisCluster() + .withBinaryCluster(conn -> conn.sync().hlen(sharedMrmKey)); + + // We expect a single device for the single resolved destination and then the data field + assertEquals(2, hashFieldCount); + } + @Test void testInsertDuplicateKey() throws Exception { final MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript insertMrmScript = new MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript( REDIS_CLUSTER_EXTENSION.getRedisCluster(), mock(ScheduledExecutorService.class)); + final SealedSenderMultiRecipientMessage multiRecipientMessage = + MessagesCacheTest.generateRandomMrmMessage(new AciServiceIdentifier(UUID.randomUUID()), Device.PRIMARY_ID); + + final Set resolvedRecipients = + new HashSet<>(multiRecipientMessage.getRecipients().values()); + final byte[] sharedMrmKey = MessagesCache.getSharedMrmKey(UUID.randomUUID()); - insertMrmScript.executeAsync(sharedMrmKey, - MessagesCacheTest.generateRandomMrmMessage(new AciServiceIdentifier(UUID.randomUUID()), Device.PRIMARY_ID)) + insertMrmScript.executeAsync(sharedMrmKey, multiRecipientMessage, resolvedRecipients) .toCompletableFuture() .join(); final CompletionException completionException = assertThrows(CompletionException.class, () -> insertMrmScript.executeAsync(sharedMrmKey, MessagesCacheTest.generateRandomMrmMessage(new AciServiceIdentifier(UUID.randomUUID()), - Device.PRIMARY_ID)).toCompletableFuture().join()); + Device.PRIMARY_ID), resolvedRecipients).toCompletableFuture().join()); assertInstanceOf(RedisException.class, completionException.getCause()); assertTrue(completionException.getCause().getMessage() diff --git a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheRemoveRecipientViewFromMrmDataScriptTest.java b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheRemoveRecipientViewFromMrmDataScriptTest.java index 0f0859662..252319cb3 100644 --- a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheRemoveRecipientViewFromMrmDataScriptTest.java +++ b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheRemoveRecipientViewFromMrmDataScriptTest.java @@ -12,6 +12,7 @@ import io.lettuce.core.cluster.SlotHash; import java.time.Duration; import java.util.ArrayList; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Objects; @@ -23,6 +24,7 @@ import org.junit.jupiter.api.extension.RegisterExtension; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.MethodSource; import org.junit.jupiter.params.provider.ValueSource; +import org.signal.libsignal.protocol.SealedSenderMultiRecipientMessage; import org.whispersystems.textsecuregcm.identity.AciServiceIdentifier; import org.whispersystems.textsecuregcm.identity.ServiceIdentifier; import org.whispersystems.textsecuregcm.redis.RedisClusterExtension; @@ -42,8 +44,11 @@ class MessagesCacheRemoveRecipientViewFromMrmDataScriptTest { final MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript insertMrmScript = new MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript( REDIS_CLUSTER_EXTENSION.getRedisCluster(), mock(ScheduledExecutorService.class)); + final SealedSenderMultiRecipientMessage multiRecipientMessage = + MessagesCacheTest.generateRandomMrmMessage(destinations); + final byte[] sharedMrmKey = MessagesCache.getSharedMrmKey(UUID.randomUUID()); - insertMrmScript.executeAsync(sharedMrmKey, MessagesCacheTest.generateRandomMrmMessage(destinations)) + insertMrmScript.executeAsync(sharedMrmKey, multiRecipientMessage, new HashSet<>(multiRecipientMessage.getRecipients().values())) .toCompletableFuture() .join(); @@ -105,9 +110,12 @@ class MessagesCacheRemoveRecipientViewFromMrmDataScriptTest { final MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript insertMrmScript = new MessagesCacheInsertSharedMultiRecipientPayloadAndViewsScript( REDIS_CLUSTER_EXTENSION.getRedisCluster(), mock(ScheduledExecutorService.class)); + final SealedSenderMultiRecipientMessage multiRecipientMessage = + MessagesCacheTest.generateRandomMrmMessage(serviceIdentifier, deviceId); + final byte[] sharedMrmKey = MessagesCache.getSharedMrmKey(UUID.randomUUID()); insertMrmScript.executeAsync(sharedMrmKey, - MessagesCacheTest.generateRandomMrmMessage(serviceIdentifier, deviceId)).toCompletableFuture().join(); + multiRecipientMessage, new HashSet<>(multiRecipientMessage.getRecipients().values())).toCompletableFuture().join(); sharedMrmKeys.add(sharedMrmKey); } diff --git a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheTest.java b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheTest.java index 0800c6299..f3affefdb 100644 --- a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheTest.java +++ b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesCacheTest.java @@ -556,7 +556,7 @@ class MessagesCacheTest { final byte[] sharedMrmDataKey; if (sharedMrmKeyPresent) { - sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm).join(); + sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm, new HashSet<>(mrm.getRecipients().values())).join(); } else { sharedMrmDataKey = "{1}".getBytes(StandardCharsets.UTF_8); } @@ -641,7 +641,7 @@ class MessagesCacheTest { .setContent(ByteString.copyFrom(mrm.messageForRecipient(recepient))) .build(); expectedQueueSize += message.getSerializedSize(); - byte[] sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm).join(); + byte[] sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm, new HashSet<>(mrm.getRecipients().values())).join(); // Insert the MRM message without the content yield message @@ -696,7 +696,7 @@ class MessagesCacheTest { final byte[] sharedMrmDataKey; if (sharedMrmKeyPresent) { - sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm).join(); + sharedMrmDataKey = messagesCache.insertSharedMultiRecipientMessagePayload(mrm, new HashSet<>(mrm.getRecipients().values())).join(); } else { sharedMrmDataKey = new byte[]{1}; } diff --git a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesManagerTest.java b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesManagerTest.java index b67714507..4dec0a5f7 100644 --- a/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesManagerTest.java +++ b/service/src/test/java/org/whispersystems/textsecuregcm/storage/MessagesManagerTest.java @@ -239,7 +239,7 @@ class MessagesManagerTest { final byte[] sharedMrmKey = "shared-mrm-key".getBytes(StandardCharsets.UTF_8); - when(messagesCache.insertSharedMultiRecipientMessagePayload(multiRecipientMessage)) + when(messagesCache.insertSharedMultiRecipientMessagePayload(eq(multiRecipientMessage), any())) .thenReturn(CompletableFuture.completedFuture(sharedMrmKey)); when(messagesCache.insert(any(), any(), anyByte(), any()))