Add a gRPC endpoint for reporting spam

This commit is contained in:
Jon Chambers authored and Jon Chambers committed 2026-09-16 10:16:40 -04:00
1 parent f66bd0858d
commit 09b5289d83
4 files changed
+113 -8

No files matched your search

@@ -1132,7 +1132,7 @@ public class WhisperServerService extends Application<WhisperServerConfiguration
new CredentialsGrpcService(accountsManager, certificateGenerator, zkAuthOperations, callingGenericZkSecretParams, rateLimiters, Clock.systemUTC(), ExternalServiceDefinitions.createExternalServiceList(config, Clock.systemUTC())),
new KeysGrpcService(accountsManager, keysManager, rateLimiters),
new ProfileGrpcService(clock, accountsManager, profilesManager, asnInfoProviderSupplier, dynamicConfigurationManager, config.getBadges(), profileCdnPolicyGenerator, chatGenericZkSecretParams, profileBadgeConverter, rateLimiters),
new MessagesGrpcService(accountsManager, rateLimiters, messageSender, messageByteLimitCardinalityEstimator, spamChecker, messageDispatcher, Clock.systemUTC()),
new MessagesGrpcService(accountsManager, reportMessageManager, phoneNumberIdentifiers, rateLimiters, messageSender, messageByteLimitCardinalityEstimator, spamChecker, messageDispatcher, Clock.systemUTC()),
new BackupsGrpcService(accountsManager, backupAuthManager, backupMetrics),
new DevicesGrpcService(accountsManager),
new AttachmentsGrpcService(experimentEnrollmentManager, rateLimiters, gcsAttachmentGenerator,
@@ -20,6 +20,8 @@ import org.signal.chat.errors.NotFound;
import org.signal.chat.messages.GetMessagesRequest;
import org.signal.chat.messages.GetMessagesResponse;
import org.signal.chat.messages.IndividualRecipientMessageBundle;
import org.signal.chat.messages.ReportMessageRequest;
import org.signal.chat.messages.ReportMessageResponse;
import org.signal.chat.messages.SendAuthenticatedSenderMessageRequest;
import org.signal.chat.messages.SendMessageAuthenticatedSenderResponse;
import org.signal.chat.messages.SendMessageType;
@@ -44,6 +46,9 @@ import org.whispersystems.textsecuregcm.spam.SpamChecker;
import org.whispersystems.textsecuregcm.storage.Account;
import org.whispersystems.textsecuregcm.storage.AccountsManager;
import org.whispersystems.textsecuregcm.storage.Device;
import org.whispersystems.textsecuregcm.storage.PhoneNumberIdentifiers;
import org.whispersystems.textsecuregcm.storage.ReportMessageHelper;
import org.whispersystems.textsecuregcm.storage.ReportMessageManager;
import org.whispersystems.textsecuregcm.util.UUIDUtil;
import reactor.adapter.JdkFlowAdapter;
import reactor.core.publisher.Flux;
@@ -52,6 +57,8 @@ import javax.annotation.Nullable;
public class MessagesGrpcService extends SimpleMessagesGrpc.MessagesImplBase {
private final AccountsManager accountsManager;
private final ReportMessageManager reportMessageManager;
private final PhoneNumberIdentifiers phoneNumberIdentifiers;
private final RateLimiters rateLimiters;
private final MessageSender messageSender;
private final CardinalityEstimator messageByteLimitEstimator;
@@ -63,6 +70,8 @@ public class MessagesGrpcService extends SimpleMessagesGrpc.MessagesImplBase {
SendMessageAuthenticatedSenderResponse.newBuilder().setSuccess(Empty.getDefaultInstance()).build();
public MessagesGrpcService(final AccountsManager accountsManager,
final ReportMessageManager reportMessageManager,
final PhoneNumberIdentifiers phoneNumberIdentifiers,
final RateLimiters rateLimiters,
final MessageSender messageSender,
final CardinalityEstimator messageByteLimitEstimator,
@@ -71,6 +80,8 @@ public class MessagesGrpcService extends SimpleMessagesGrpc.MessagesImplBase {
final Clock clock) {
this.accountsManager = accountsManager;
this.reportMessageManager = reportMessageManager;
this.phoneNumberIdentifiers = phoneNumberIdentifiers;
this.rateLimiters = rateLimiters;
this.messageSender = messageSender;
this.messageByteLimitEstimator = messageByteLimitEstimator;
@@ -261,4 +272,22 @@ public class MessagesGrpcService extends SimpleMessagesGrpc.MessagesImplBase {
throw GrpcExceptions.invalidArguments("unrecognized envelope type");
};
}
@Override
public ReportMessageResponse reportMessage(final ReportMessageRequest request) {
final ServiceIdentifier sourceServiceIdentifier =
GrpcServiceIdentifierUtil.fromGrpcServiceIdentifier(request.getSourceServiceIdentifier());
ReportMessageHelper.reportMessage(
sourceServiceIdentifier,
new AciServiceIdentifier(AuthenticationUtil.requireAuthenticatedDevice().accountIdentifier()),
UUIDUtil.fromByteString(request.getMessageGuid()),
request.getReportSpamToken().isEmpty() ? null : request.getReportSpamToken().toByteArray(),
RequestAttributesUtil.getUserAgent().orElse(null),
accountsManager,
phoneNumberIdentifiers,
reportMessageManager);
return ReportMessageResponse.getDefaultInstance();
}
}
@@ -49,6 +49,9 @@ service Messages {
// a STREAM_CLOSED error reason. A GetMessagesStreamClosed message will be
// present in the error details.
rpc GetMessages(stream GetMessagesRequest) returns (stream GetMessagesResponse) {}
// Reports a message as spam.
rpc ReportMessage(ReportMessageRequest) returns (ReportMessageResponse) {}
}
message GetMessagesRequest {
@@ -435,3 +438,18 @@ message ChallengeRequired {
// resolved by waiting.
optional uint64 retry_after_seconds = 3;
}
message ReportMessageRequest {
// The service identifier of the party that sent the offending message
common.ServiceIdentifier source_service_identifier = 1;
// The GUID of the offending message
bytes message_guid = 2 [(require.exactlySize) = 16];
// The spam-reporting token attached to the offending message
bytes report_spam_token = 3;
}
message ReportMessageResponse {
}
@@ -12,6 +12,7 @@ import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyBoolean;
import static org.mockito.ArgumentMatchers.anyByte;
import static org.mockito.ArgumentMatchers.anyLong;
import static org.mockito.ArgumentMatchers.argThat;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.doAnswer;
import static org.mockito.Mockito.doThrow;
@@ -30,6 +31,7 @@ import io.grpc.StatusRuntimeException;
import io.grpc.stub.BlockingClientCall;
import java.time.Duration;
import java.time.Instant;
import java.util.Arrays;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
@@ -38,7 +40,6 @@ import java.util.Set;
import java.util.UUID;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.stream.Stream;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Nested;
@@ -56,6 +57,7 @@ import org.signal.chat.messages.GetMessagesResponse;
import org.signal.chat.messages.IndividualRecipientMessageBundle;
import org.signal.chat.messages.MessagesGrpc;
import org.signal.chat.messages.MismatchedDevices;
import org.signal.chat.messages.ReportMessageRequest;
import org.signal.chat.messages.SendAuthenticatedSenderMessageRequest;
import org.signal.chat.messages.SendMessageAuthenticatedSenderResponse;
import org.signal.chat.messages.SendMessageType;
@@ -66,7 +68,7 @@ import org.whispersystems.textsecuregcm.controllers.MismatchedDevicesException;
import org.whispersystems.textsecuregcm.controllers.RateLimitExceededException;
import org.whispersystems.textsecuregcm.entities.MessageProtos;
import org.whispersystems.textsecuregcm.identity.AciServiceIdentifier;
import org.whispersystems.textsecuregcm.identity.IdentityType;
import org.whispersystems.textsecuregcm.identity.PniServiceIdentifier;
import org.whispersystems.textsecuregcm.identity.ServiceIdentifier;
import org.whispersystems.textsecuregcm.limits.CardinalityEstimator;
import org.whispersystems.textsecuregcm.limits.RateLimiter;
@@ -80,6 +82,8 @@ import org.whispersystems.textsecuregcm.spam.SpamChecker;
import org.whispersystems.textsecuregcm.storage.Account;
import org.whispersystems.textsecuregcm.storage.AccountsManager;
import org.whispersystems.textsecuregcm.storage.Device;
import org.whispersystems.textsecuregcm.storage.PhoneNumberIdentifiers;
import org.whispersystems.textsecuregcm.storage.ReportMessageManager;
import org.whispersystems.textsecuregcm.tests.util.DevicesHelper;
import org.whispersystems.textsecuregcm.util.TestClock;
import org.whispersystems.textsecuregcm.util.TestRandomUtil;
@@ -92,6 +96,12 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
@Mock
private AccountsManager accountsManager;
@Mock
private ReportMessageManager reportMessageManager;
@Mock
private PhoneNumberIdentifiers phoneNumberIdentifiers;
@Mock
private RateLimiters rateLimiters;
@@ -140,6 +150,8 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
@Override
protected MessagesGrpcService createServiceBeforeEachTest() {
return new MessagesGrpcService(accountsManager,
reportMessageManager,
phoneNumberIdentifiers,
rateLimiters,
messageSender,
messageByteLimitEstimator,
@@ -378,7 +390,6 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
.setType(SendMessageType.DOUBLE_RATCHET)
.build());
//noinspection ResultOfMethodCallIgnored
assertRateLimitExceeded(retryDuration,
() -> authenticatedServiceStub().sendMessage(
generateRequest(serviceIdentifier, false, true, messages)));
@@ -682,7 +693,6 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
.setType(SendMessageType.DOUBLE_RATCHET)
.build());
//noinspection ResultOfMethodCallIgnored
assertRateLimitExceeded(retryDuration, () ->
authenticatedServiceStub().sendSyncMessage(generateRequest(true, messages)));
@@ -719,7 +729,6 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
@Test
void messageDeliveryNotAllowed()
throws MessageTooLargeException, MessageDeliveryNotAllowedException, MismatchedDevicesException {
final AciServiceIdentifier serviceIdentifier = new AciServiceIdentifier(AUTHENTICATED_ACI);
final byte[] payload = TestRandomUtil.nextBytes(128);
final Map<Byte, IndividualRecipientMessageBundle.Message> messages =
@@ -783,7 +792,7 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
@ParameterizedTest
@MethodSource
void invalidAckMessages(final GetMessagesRequest request)
throws StatusException, InterruptedException, TimeoutException {
throws StatusException, InterruptedException {
doAnswer(invocation -> {
Flux<UUID> ackArg = invocation.getArgument(4);
// use mapNotNull instead of `then` because there is an interaction between the blocking client and
@@ -793,7 +802,7 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
}).when(messageDispatcher).getMessages(anyBoolean(), any(), any(), any(), any());
final BlockingClientCall<GetMessagesRequest, GetMessagesResponse> blockingCall = authenticatedServiceStub().getMessages();
final CompletableFuture<StatusException> reader = CompletableFuture.supplyAsync(() -> assertThrows(StatusException.class, () -> blockingCall.read()));
final CompletableFuture<StatusException> reader = CompletableFuture.supplyAsync(() -> assertThrows(StatusException.class, blockingCall::read));
blockingCall.write(GetMessagesRequest.newBuilder().setOptions(GetMessagesRequest.GetMessageOptions.getDefaultInstance()).build());
blockingCall.write(request);
@@ -814,6 +823,55 @@ class MessagesGrpcServiceTest extends SimpleBaseGrpcTest<MessagesGrpcService, Me
}
}
@Timeout(value = 1, unit = TimeUnit.MINUTES, threadMode = Timeout.ThreadMode.SEPARATE_THREAD)
@Nested
class ReportMessage {
@Test
void reportMessage() throws StatusException {
final AciServiceIdentifier sourceServiceIdentifier = new AciServiceIdentifier(UUID.randomUUID());
final UUID messageGuid = UUID.randomUUID();
final byte[] reportSpamToken = TestRandomUtil.nextBytes(128);
//noinspection ResultOfMethodCallIgnored
authenticatedServiceStub().reportMessage(ReportMessageRequest.newBuilder()
.setSourceServiceIdentifier(GrpcServiceIdentifierUtil.toGrpcServiceIdentifier(sourceServiceIdentifier))
.setMessageGuid(UUIDUtil.toByteString(messageGuid))
.setReportSpamToken(ByteString.copyFrom(reportSpamToken))
.build());
verify(reportMessageManager).report(eq(Optional.empty()),
eq(sourceServiceIdentifier.uuid()),
eq(Optional.empty()),
eq(messageGuid),
eq(AUTHENTICATED_ACI),
argThat(maybeToken -> maybeToken.map(token -> Arrays.equals(token, reportSpamToken)).orElse(false)),
any(),
eq(true));
}
@Test
void reportMessageNoToken() throws StatusException {
final AciServiceIdentifier sourceServiceIdentifier = new AciServiceIdentifier(UUID.randomUUID());
final UUID messageGuid = UUID.randomUUID();
//noinspection ResultOfMethodCallIgnored
authenticatedServiceStub().reportMessage(ReportMessageRequest.newBuilder()
.setSourceServiceIdentifier(GrpcServiceIdentifierUtil.toGrpcServiceIdentifier(sourceServiceIdentifier))
.setMessageGuid(UUIDUtil.toByteString(messageGuid))
.build());
verify(reportMessageManager).report(eq(Optional.empty()),
eq(sourceServiceIdentifier.uuid()),
eq(Optional.empty()),
eq(messageGuid),
eq(AUTHENTICATED_ACI),
eq(Optional.empty()),
any(),
eq(true));
}
}
private static ThrowingSupplier<?> convertStatusException(final ThrowingSupplier<?> serviceCall) {
return () -> {
try {