mirror of
https://github.com/signalapp/Signal-Server
synced 2026-10-05 20:47:55 +01:00
Don't insert shared MRM payloads for unresolved recipients
This commit is contained in:
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));
|
||||
|
||||
+19
-5
@@ -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));
|
||||
}
|
||||
});
|
||||
|
||||
+22
-13
@@ -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).
|
||||
|
||||
+49
-5
@@ -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()
|
||||
|
||||
+10
-2
@@ -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);
|
||||
}
|
||||
|
||||
+3
-3
@@ -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};
|
||||
}
|
||||
|
||||
+1
-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()))
|
||||
|
||||
Reference in new issue
Block a user