Add an integration test for recovering with a TOTP

This commit is contained in:
Ravi Khadiwala authored and ravi-signal committed 2026-09-09 16:31:03 -05:00
1 parent 53bd7a5587
commit 5221b3fb84
5 files changed
+173 -274

No files matched your search

@@ -1,101 +0,0 @@
/*
* Copyright 2023 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.signal.integration;
import com.fasterxml.jackson.core.JsonGenerator;
import com.fasterxml.jackson.core.JsonParser;
import com.fasterxml.jackson.databind.DeserializationContext;
import com.fasterxml.jackson.databind.JsonDeserializer;
import com.fasterxml.jackson.databind.JsonSerializer;
import com.fasterxml.jackson.databind.SerializerProvider;
import java.io.IOException;
import java.util.Base64;
import org.signal.libsignal.protocol.IdentityKey;
import org.signal.libsignal.protocol.ecc.ECPublicKey;
public final class Codecs {
private Codecs() {
// utility class
}
@FunctionalInterface
public interface CheckedFunction<T, R> {
R apply(T t) throws Exception;
}
public static class Base64BasedSerializer<T> extends JsonSerializer<T> {
private final CheckedFunction<T, byte[]> mapper;
public Base64BasedSerializer(final CheckedFunction<T, byte[]> mapper) {
this.mapper = mapper;
}
@Override
public void serialize(final T value, final JsonGenerator gen, final SerializerProvider serializers) throws IOException {
try {
gen.writeString(Base64.getEncoder().withoutPadding().encodeToString(mapper.apply(value)));
} catch (Exception e) {
throw new RuntimeException(e);
}
}
}
public static class Base64BasedDeserializer<T> extends JsonDeserializer<T> {
private final CheckedFunction<byte[], T> mapper;
public Base64BasedDeserializer(final CheckedFunction<byte[], T> mapper) {
this.mapper = mapper;
}
@Override
public T deserialize(final JsonParser p, final DeserializationContext ctxt) throws IOException {
try {
return mapper.apply(Base64.getDecoder().decode(p.getValueAsString()));
} catch (Exception e) {
throw new RuntimeException(e);
}
}
}
public static class ByteArraySerializer extends Base64BasedSerializer<byte[]> {
public ByteArraySerializer() {
super(bytes -> bytes);
}
}
public static class ByteArrayDeserializer extends Base64BasedDeserializer<byte[]> {
public ByteArrayDeserializer() {
super(bytes -> bytes);
}
}
public static class ECPublicKeySerializer extends Base64BasedSerializer<ECPublicKey> {
public ECPublicKeySerializer() {
super(ECPublicKey::serialize);
}
}
public static class ECPublicKeyDeserializer extends Base64BasedDeserializer<ECPublicKey> {
public ECPublicKeyDeserializer() {
super(ECPublicKey::new);
}
}
public static class IdentityKeySerializer extends Base64BasedSerializer<IdentityKey> {
public IdentityKeySerializer() {
super(IdentityKey::serialize);
}
}
public static class IdentityKeyDeserializer extends Base64BasedDeserializer<IdentityKey> {
public IdentityKeyDeserializer() {
super(bytes -> new IdentityKey(bytes, 0));
}
}
}
@@ -12,8 +12,17 @@ import com.google.common.io.Resources;
import com.google.common.net.HttpHeaders;
import io.dropwizard.configuration.ConfigurationValidationException;
import io.dropwizard.jersey.validation.Validators;
import io.grpc.ChannelCredentials;
import io.grpc.ClientInterceptor;
import io.grpc.Grpc;
import io.grpc.ManagedChannel;
import io.grpc.Metadata;
import io.grpc.TlsChannelCredentials;
import io.grpc.stub.MetadataUtils;
import jakarta.validation.ConstraintViolation;
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.UncheckedIOException;
import java.lang.invoke.MethodHandles;
import java.net.URI;
import java.net.URL;
@@ -32,6 +41,7 @@ import java.util.Optional;
import java.util.Set;
import java.util.concurrent.Executors;
import java.util.concurrent.TimeUnit;
import javax.annotation.Nullable;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.lang3.Validate;
import org.apache.commons.lang3.tuple.Pair;
@@ -58,6 +68,7 @@ import org.signal.libsignal.zkgroup.receipts.ReceiptCredentialPresentation;
import org.signal.libsignal.zkgroup.receipts.ReceiptSerial;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.whispersystems.textsecuregcm.auth.grpc.RequireAuthenticationInterceptor;
import org.whispersystems.textsecuregcm.entities.AccountAttributes;
import org.whispersystems.textsecuregcm.entities.AccountIdentityResponse;
import org.whispersystems.textsecuregcm.entities.DeviceActivationRequest;
@@ -77,6 +88,8 @@ public final class Operations {
private static final Config CONFIG = loadConfigFromClasspath("config.yml");
private static final String GRPC_DOMAIN = "grpc." + CONFIG.domain();
private static final IntegrationTools INTEGRATION_TOOLS = IntegrationTools.create(CONFIG);
private static final String USER_AGENT = "integration-test";
@@ -85,6 +98,7 @@ public final class Operations {
private static final WebSocketClient WEB_SOCKET_CLIENT = buildWebSocketClient();
private static final ManagedChannel GRPC_CHANNEL = buildGrpcChannel();
private Operations() {
// utility class
@@ -133,6 +147,36 @@ public final class Operations {
return user;
}
public static TestUser recoverNumberlessUser(final TestUser testUser, @Nullable final Integer totp) {
final String accountPassword = Base64.getEncoder().encodeToString(randomBytes(32));
final TestUser recoveredUser = TestUser.createNumberlessForRecovery(accountPassword, testUser.registrationPassword());
final ECKeyPair aciIdentityKeyPair = ECKeyPair.generate();
final ECKeyPair pniIdentityKeyPair = ECKeyPair.generate();
final RegistrationRequest registrationRequest = new RegistrationRequest(null,
testUser.registrationPassword(),
null,
totp,
recoveredUser.accountAttributes(),
true,
new IdentityKey(aciIdentityKeyPair.getPublicKey()),
new IdentityKey(pniIdentityKeyPair.getPublicKey()),
new DeviceActivationRequest(generateSignedECPreKey(1, aciIdentityKeyPair),
Optional.of(generateSignedECPreKey(2, pniIdentityKeyPair)),
generateSignedKEMPreKey(3, aciIdentityKeyPair),
Optional.of(generateSignedKEMPreKey(4, pniIdentityKeyPair)),
Optional.empty(),
Optional.empty()));
final AccountIdentityResponse registrationResponse = apiPost("/v1/registration", registrationRequest)
// For a numberless account recovery, the username is the ACI
.authorized(testUser.aciUuid().toString(), accountPassword)
.executeExpectSuccess(AccountIdentityResponse.class);
recoveredUser.setAciUuid(registrationResponse.uuid());
return recoveredUser;
}
public static TestUser newRegisteredUser(final String number) {
final byte[] registrationPassword = populateRandomRecoveryPassword(number);
final String accountPassword = Base64.getEncoder().encodeToString(randomBytes(32));
@@ -239,7 +283,7 @@ public final class Operations {
}
}
private static byte[] randomBytes(int numBytes) {
public static byte[] randomBytes(final int numBytes) {
final byte[] bytes = new byte[numBytes];
new SecureRandom().nextBytes(bytes);
return bytes;
@@ -272,6 +316,33 @@ public final class Operations {
return URI.create("https://" + CONFIG.domain() + endpoint + query);
}
public static ManagedChannel grpcChannel() {
return GRPC_CHANNEL;
}
public static ClientInterceptor authorizationInterceptor(final TestUser user, final byte deviceId) {
final String username = "%s.%d".formatted(user.aciUuid().toString(), deviceId);
final Metadata metadata = new Metadata();
metadata.put(RequireAuthenticationInterceptor.AUTHORIZATION_METADATA_KEY,
HeaderUtils.basicAuthHeader(username, user.accountPassword()));
return MetadataUtils.newAttachHeadersInterceptor(metadata);
}
private static ManagedChannel buildGrpcChannel() {
try {
final ByteArrayInputStream rootCert =
new ByteArrayInputStream(CONFIG.rootCert().getBytes(StandardCharsets.UTF_8));
final ChannelCredentials credentials = TlsChannelCredentials.newBuilder().trustManager(rootCert).build();
return Grpc.newChannelBuilderForAddress(GRPC_DOMAIN, 443, credentials)
.userAgent(USER_AGENT)
.build();
} catch (final IOException e) {
throw new UncheckedIOException(e);
}
}
public static class RequestBuilder {
private final HttpRequest.Builder builder;
@@ -416,7 +487,7 @@ public final class Operations {
final String path,
final Map<String, String> headers) throws IOException {
final URI uri = URI.create("wss://grpc." + CONFIG.domain() + path);
final URI uri = URI.create("wss://" + GRPC_DOMAIN + path);
final ClientUpgradeRequest request = new ClientUpgradeRequest(uri);
headers.forEach(request::setHeader);
@@ -1,61 +0,0 @@
/*
* Copyright 2023 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.signal.integration;
import jakarta.annotation.Nullable;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import org.apache.commons.lang3.tuple.Pair;
import org.signal.libsignal.protocol.IdentityKeyPair;
import org.signal.libsignal.protocol.ecc.ECKeyPair;
import org.signal.libsignal.protocol.state.SignedPreKeyRecord;
public class TestDevice {
private final byte deviceId;
private final Map<Integer, Pair<IdentityKeyPair, SignedPreKeyRecord>> signedPreKeys = new ConcurrentHashMap<>();
public static TestDevice create(
final byte deviceId,
final IdentityKeyPair aciIdentityKeyPair,
@Nullable final IdentityKeyPair pniIdentityKeyPair) {
final TestDevice device = new TestDevice(deviceId);
device.addSignedPreKey(aciIdentityKeyPair);
if (pniIdentityKeyPair != null) {
device.addSignedPreKey(pniIdentityKeyPair);
}
return device;
}
public TestDevice(final byte deviceId) {
this.deviceId = deviceId;
}
public byte deviceId() {
return deviceId;
}
public SignedPreKeyRecord latestSignedPreKey(final IdentityKeyPair identity) {
final int id = signedPreKeys.entrySet()
.stream()
.filter(p -> p.getValue().getLeft().equals(identity))
.mapToInt(Map.Entry::getKey)
.max()
.orElseThrow();
return signedPreKeys.get(id).getRight();
}
public SignedPreKeyRecord addSignedPreKey(final IdentityKeyPair identity) {
final int nextId = signedPreKeys.keySet().stream().mapToInt(k -> k + 1).max().orElse(0);
final ECKeyPair keyPair = ECKeyPair.generate();
final byte[] signature = keyPair.getPrivateKey().calculateSignature(keyPair.getPublicKey().serialize());
final SignedPreKeyRecord signedPreKeyRecord = new SignedPreKeyRecord(nextId, System.currentTimeMillis(), keyPair, signature);
signedPreKeys.put(nextId, Pair.of(identity, signedPreKeyRecord));
return signedPreKeyRecord;
}
}
@@ -5,29 +5,15 @@
package org.signal.integration;
import static java.util.Objects.requireNonNull;
import com.fasterxml.jackson.databind.annotation.JsonDeserialize;
import com.fasterxml.jackson.databind.annotation.JsonSerialize;
import java.security.SecureRandom;
import java.nio.charset.StandardCharsets;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.security.SecureRandom;
import java.util.Optional;
import java.util.UUID;
import java.util.concurrent.ConcurrentHashMap;
import org.signal.libsignal.protocol.IdentityKey;
import org.signal.libsignal.protocol.IdentityKeyPair;
import org.signal.libsignal.protocol.InvalidKeyException;
import org.signal.libsignal.protocol.ecc.ECPublicKey;
import org.signal.libsignal.protocol.state.SignedPreKeyRecord;
import javax.annotation.Nullable;
import org.signal.libsignal.protocol.util.KeyHelper;
import org.whispersystems.textsecuregcm.auth.UnidentifiedAccessUtil;
import org.whispersystems.textsecuregcm.entities.AccountAttributes;
import org.whispersystems.textsecuregcm.storage.Device;
import org.whispersystems.textsecuregcm.storage.DeviceCapability;
import javax.annotation.Nullable;
public class TestUser {
@@ -36,21 +22,14 @@ public class TestUser {
@Nullable
private final Integer pniRegistrationId;
private final IdentityKeyPair aciIdentityKey;
private final Map<Byte, TestDevice> devices = new ConcurrentHashMap<>();
private final byte[] unidentifiedAccessKey;
@Nullable
private String phoneNumber;
private final String phoneNumber;
@Nullable
private IdentityKeyPair pniIdentityKey;
private final String accountPassword;
private String accountPassword;
private byte[] registrationPassword;
private final byte[] registrationPassword;
private UUID aciUuid;
@@ -58,40 +37,44 @@ public class TestUser {
private UUID pniUuid;
public static TestUser createNumberless(final String accountPassword, final byte[] accountRecoveryPassword) {
final IdentityKeyPair aciIdentityKey = IdentityKeyPair.generate();
final int registrationId = KeyHelper.generateRegistrationId(false);
final byte[] unidentifiedAccessKey = new byte[UnidentifiedAccessUtil.UNIDENTIFIED_ACCESS_KEY_LENGTH];
new SecureRandom().nextBytes(unidentifiedAccessKey);
final byte[] unidentifiedAccessKey = Operations.randomBytes(UnidentifiedAccessUtil.UNIDENTIFIED_ACCESS_KEY_LENGTH);
return new TestUser(
registrationId,
null,
aciIdentityKey,
null,
null,
unidentifiedAccessKey,
accountPassword,
accountRecoveryPassword);
}
public static TestUser createNumberlessForRecovery(final String accountPassword, final byte[] accountRecoveryPassword) {
// Recovering a numberless account requires PNI keys (though they are discarded by the server)
final int registrationId = KeyHelper.generateRegistrationId(false);
final int pniRegistrationId = KeyHelper.generateRegistrationId(false);
final byte[] unidentifiedAccessKey = new byte[UnidentifiedAccessUtil.UNIDENTIFIED_ACCESS_KEY_LENGTH];
new SecureRandom().nextBytes(unidentifiedAccessKey);
return new TestUser(
registrationId,
pniRegistrationId,
null,
unidentifiedAccessKey,
accountPassword,
accountRecoveryPassword);
}
public static TestUser create(final String phoneNumber, final String accountPassword, final byte[] registrationPassword) {
// ACI identity key pair
final IdentityKeyPair aciIdentityKey = IdentityKeyPair.generate();
// PNI identity key pair
final IdentityKeyPair pniIdentityKey = IdentityKeyPair.generate();
// registration id
final int registrationId = KeyHelper.generateRegistrationId(false);
final int pniRegistrationId = KeyHelper.generateRegistrationId(false);
// uak
final byte[] unidentifiedAccessKey = new byte[UnidentifiedAccessUtil.UNIDENTIFIED_ACCESS_KEY_LENGTH];
new SecureRandom().nextBytes(unidentifiedAccessKey);
return new TestUser(
registrationId,
pniRegistrationId,
aciIdentityKey,
phoneNumber,
pniIdentityKey,
unidentifiedAccessKey,
accountPassword,
registrationPassword);
@@ -100,39 +83,26 @@ public class TestUser {
public TestUser(
final int registrationId,
@Nullable final Integer pniRegistrationId,
final IdentityKeyPair aciIdentityKey,
@Nullable final String phoneNumber,
@Nullable final IdentityKeyPair pniIdentityKey,
final byte[] unidentifiedAccessKey,
final String accountPassword,
final byte[] registrationPassword) {
this.registrationId = registrationId;
this.pniRegistrationId = pniRegistrationId;
this.aciIdentityKey = aciIdentityKey;
this.phoneNumber = phoneNumber;
this.pniIdentityKey = pniIdentityKey;
this.unidentifiedAccessKey = unidentifiedAccessKey;
this.accountPassword = accountPassword;
this.registrationPassword = registrationPassword;
devices.put(Device.PRIMARY_ID, TestDevice.create(Device.PRIMARY_ID, aciIdentityKey, pniIdentityKey));
}
public int registrationId() {
return registrationId;
}
public IdentityKeyPair aciIdentityKey() {
return aciIdentityKey;
}
public Optional<String> phoneNumber() {
return Optional.ofNullable(phoneNumber);
}
public Optional<IdentityKeyPair> pniIdentityKey() {
return Optional.ofNullable(pniIdentityKey);
}
public String accountPassword() {
return accountPassword;
}
@@ -151,9 +121,8 @@ public class TestUser {
public AccountAttributes accountAttributes() {
return new AccountAttributes(true, registrationId, pniRegistrationId, "".getBytes(StandardCharsets.UTF_8), "", true,
DeviceCapability.CAPABILITIES_REQUIRED_FOR_NEW_DEVICES, null)
.setUnidentifiedAccessKey(unidentifiedAccessKey)
.setRecoveryPassword(registrationPassword);
DeviceCapability.CAPABILITIES_REQUIRED_FOR_NEW_DEVICES, registrationPassword)
.setUnidentifiedAccessKey(unidentifiedAccessKey);
}
public void setAciUuid(final UUID aciUuid) {
@@ -163,59 +132,4 @@ public class TestUser {
public void setPniUuid(@Nullable final UUID pniUuid) {
this.pniUuid = pniUuid;
}
public void setPhoneNumber(@Nullable final String phoneNumber) {
this.phoneNumber = phoneNumber;
}
public void setPniIdentityKey(@Nullable final IdentityKeyPair pniIdentityKey) {
this.pniIdentityKey = pniIdentityKey;
}
public void setAccountPassword(final String accountPassword) {
this.accountPassword = accountPassword;
}
public void setRegistrationPassword(final byte[] registrationPassword) {
this.registrationPassword = registrationPassword;
}
public PreKeySetPublicView preKeys(final byte deviceId, final boolean pni) {
final IdentityKeyPair identity = pni
? pniIdentityKey
: aciIdentityKey;
final TestDevice device = requireNonNull(devices.get(deviceId));
final SignedPreKeyRecord signedPreKeyRecord = device.latestSignedPreKey(identity);
try {
return new PreKeySetPublicView(
Collections.emptyList(),
identity.getPublicKey(),
new SignedPreKeyPublicView(
signedPreKeyRecord.getId(),
signedPreKeyRecord.getKeyPair().getPublicKey(),
signedPreKeyRecord.getSignature()
)
);
} catch (InvalidKeyException e) {
throw new RuntimeException(e);
}
}
public record SignedPreKeyPublicView(
int keyId,
@JsonSerialize(using = Codecs.ECPublicKeySerializer.class)
@JsonDeserialize(using = Codecs.ECPublicKeyDeserializer.class)
ECPublicKey publicKey,
@JsonSerialize(using = Codecs.ByteArraySerializer.class)
@JsonDeserialize(using = Codecs.ByteArrayDeserializer.class)
byte[] signature) {
}
public record PreKeySetPublicView(
List<String> preKeys,
@JsonSerialize(using = Codecs.IdentityKeySerializer.class)
@JsonDeserialize(using = Codecs.IdentityKeyDeserializer.class)
IdentityKey identityKey,
SignedPreKeyPublicView signedPreKey) {
}
}
@@ -5,18 +5,37 @@
package org.signal.integration;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import com.eatthepath.otp.TimeBasedOneTimePasswordGenerator;
import com.google.protobuf.ByteString;
import java.security.InvalidKeyException;
import java.security.NoSuchAlgorithmException;
import java.time.Duration;
import java.time.Instant;
import java.util.Arrays;
import java.util.Base64;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.UUID;
import javax.crypto.SecretKey;
import javax.crypto.spec.SecretKeySpec;
import org.apache.commons.lang3.tuple.Pair;
import org.apache.http.HttpStatus;
import org.junit.jupiter.api.Test;
import org.signal.chat.account.AccountsGrpc;
import org.signal.chat.account.ConfirmTotpKeyRequest;
import org.signal.chat.account.ConfirmTotpKeyResponse;
import org.signal.chat.account.GenerateTotpKeyRequest;
import org.signal.chat.account.GenerateTotpKeyResponse;
import org.signal.chat.account.ListMfaKeysRequest;
import org.signal.chat.account.ListMfaKeysResponse;
import org.signal.chat.account.TotpParameters;
import org.signal.libsignal.protocol.IdentityKey;
import org.signal.libsignal.protocol.ecc.ECKeyPair;
import org.signal.libsignal.usernames.BaseUsernameException;
@@ -51,6 +70,63 @@ public class AccountTest {
}
}
@Test
public void testRecoverWithTotp()
throws VerificationFailedException, InvalidInputException, NoSuchAlgorithmException, InvalidKeyException {
final Operations.Receipt receipt = Operations.getPrescribedReceipt();
final TestUser originalUser = Operations.registerNumberlessUser(receipt.credential());
try {
final GenerateTotpKeyResponse generateTotpKeyResponse =
getAccountsStubForUser(originalUser).generateTotpKey(GenerateTotpKeyRequest.getDefaultInstance());
assertEquals(GenerateTotpKeyResponse.ResponseCase.KEY_GENERATED, generateTotpKeyResponse.getResponseCase());
final TotpParameters totpParameters = generateTotpKeyResponse.getKeyGenerated().getTotpParameters();
final TimeBasedOneTimePasswordGenerator totpGenerator = new TimeBasedOneTimePasswordGenerator(
Duration.ofSeconds(totpParameters.getTimeStepSeconds()),
totpParameters.getPasswordLength(),
totpParameters.getAlgorithm());
final SecretKey totpKey = new SecretKeySpec(
generateTotpKeyResponse.getKeyGenerated().getKey().toByteArray(),
totpParameters.getAlgorithm());
final byte[] totpMetadata = Operations.randomBytes(160);
final ConfirmTotpKeyResponse confirmTotpKeyResponse = getAccountsStubForUser(originalUser)
.confirmTotpKey(ConfirmTotpKeyRequest.newBuilder()
.setOneTimePassword(totpGenerator.generateOneTimePassword(totpKey, Instant.now()))
.setMetadataCiphertext(ByteString.copyFrom(totpMetadata))
.build());
assertEquals(ConfirmTotpKeyResponse.ResponseCase.KEY_CONFIRMED, confirmTotpKeyResponse.getResponseCase());
final int keyId = confirmTotpKeyResponse.getKeyConfirmed().getKeyId();
final TestUser recoveredUser =
Operations.recoverNumberlessUser(originalUser, totpGenerator.generateOneTimePassword(totpKey, Instant.now()));
assertEquals(originalUser.aciUuid(), recoveredUser.aciUuid());
// MFA key should remain set after re-registration
final Map<Integer, ListMfaKeysResponse.MfaKeyMetadata> mfaKeys =
getAccountsStubForUser(recoveredUser).listMfaKeys(ListMfaKeysRequest.getDefaultInstance()).getKeysMap();
assertEquals(1, mfaKeys.size());
assertTrue(mfaKeys.containsKey(keyId));
assertEquals(ListMfaKeysResponse.MfaKeyMetadata.MfaKeyType.MFA_KEY_TYPE_TOTP, mfaKeys.get(keyId).getType());
assertArrayEquals(totpMetadata, mfaKeys.get(keyId).getMetadataCiphertext().toByteArray());
} finally {
Operations.deleteReceipt(receipt.serial());
Operations.deleteUser(originalUser);
}
}
private static AccountsGrpc.AccountsBlockingStub getAccountsStubForUser(final TestUser user) {
return AccountsGrpc
.newBlockingStub(Operations.grpcChannel())
.withInterceptors(Operations.authorizationInterceptor(user, Device.PRIMARY_ID));
}
@Test
public void testCreateAccount() {
final TestUser user = Operations.newRegisteredUser("+19995550101");