Don't insert shared MRM payloads for unresolved recipients

This commit is contained in:
Jon Chambers authored and Chris Eager committed 2026-09-25 10:48:59 -04:00
1 parent 438ad9fb8f
commit 740116e1fd
7 files changed
+108 -31

No files matched your search

@@ -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<byte[]> insertSharedMultiRecipientMessagePayload(
final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage) {
final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage,
final Set<SealedSenderMultiRecipientMessage.Recipient> 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));
@@ -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<Void> executeAsync(final byte[] sharedMrmKey, final SealedSenderMultiRecipientMessage message) {
CompletionStage<Void> executeAsync(final byte[] sharedMrmKey,
final SealedSenderMultiRecipientMessage message,
final Set<SealedSenderMultiRecipientMessage.Recipient> resolvedRecipients) {
final List<byte[]> keys = List.of(
sharedMrmKey // sharedMrmKey
);
// Pre-allocate capacity for the most fields we expect -- 6 devices per recipient, plus the data field.
final List<byte[]> 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<byte[]> 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));
}
});
@@ -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<byte[]> 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<Optional<byte[]>> insertSharedMultiRecipientMessagePayload(
final SealedSenderMultiRecipientMessage sealedSenderMultiRecipientMessage,
final Set<SealedSenderMultiRecipientMessage.Recipient> 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).
@@ -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<ServiceIdentifier, List<Byte>> destinations = Map.of(
resolvedServiceIdentifier, List.of(Device.PRIMARY_ID),
unresolvedServiceIdentifier, List.of(Device.PRIMARY_ID));
final SealedSenderMultiRecipientMessage multiRecipientMessage =
MessagesCacheTest.generateRandomMrmMessage(destinations);
final Set<SealedSenderMultiRecipientMessage.Recipient> 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<SealedSenderMultiRecipientMessage.Recipient> 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()
@@ -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);
}
@@ -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};
}
@@ -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()))