From 034e50fedf6843cc05f212695c15d5aaaa4f2d0e Mon Sep 17 00:00:00 2001 From: Greyson Parrelli Date: Wed, 26 Aug 2026 12:11:26 -0400 Subject: [PATCH] Fix temp registration data being committed pre-registration. --- .../registration/RegistrationRepository.kt | 14 ++++- .../PinEntryForRegistrationLockViewModel.kt | 2 +- .../pinentry/PinEntryForSmsBypassViewModel.kt | 2 +- .../PinEntryForSvrRestoreViewModel.kt | 2 +- .../RegistrationRepositoryTest.kt | 60 ++++++++++++++++++- ...inEntryForRegistrationLockViewModelTest.kt | 57 +++++++++++------- .../PinEntryForSmsBypassViewModelTest.kt | 45 +++++++++----- .../PinEntryForSvrRestoreViewModelTest.kt | 33 +++++++--- 8 files changed, 167 insertions(+), 48 deletions(-) diff --git a/feature/registration/src/main/java/org/signal/registration/RegistrationRepository.kt b/feature/registration/src/main/java/org/signal/registration/RegistrationRepository.kt index f213c97037..c1038758e6 100644 --- a/feature/registration/src/main/java/org/signal/registration/RegistrationRepository.kt +++ b/feature/registration/src/main/java/org/signal/registration/RegistrationRepository.kt @@ -242,10 +242,15 @@ class RegistrationRepository( networkController.checkSvrCredentials(e164, credentials) } + /** + * @param isRegistered Whether the account is already registered. This can be called before registration (e.g. to unlock a reglocked account), and in that + * case we must not commit the restored data to persistent storage yet. + */ suspend fun restoreMasterKeyFromSvr( svrCredentials: SvrCredentials, pin: String, - forRegistrationLock: Boolean + forRegistrationLock: Boolean, + isRegistered: Boolean ): RequestResult = withContext(Dispatchers.IO) { networkController.restoreMasterKeyFromSvr( svrCredentials = svrCredentials, @@ -258,7 +263,12 @@ class RegistrationRepository( this.registrationLockEnabled = forRegistrationLock this.svrCredentials += SvrCredential(username = svrCredentials.username, password = svrCredentials.password) } - storageController.commitRegistrationData() + + if (isRegistered) { + storageController.commitRegistrationData() + } else { + Log.i(TAG, "[restoreMasterKeyFromSvr] Not yet registered. Skipping commit of registration data.") + } } } } diff --git a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModel.kt b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModel.kt index 3f6b858474..84e4f276ae 100644 --- a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModel.kt +++ b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModel.kt @@ -103,7 +103,7 @@ class PinEntryForRegistrationLockViewModel( private suspend fun applyPinEntered(state: PinEntryState, event: PinEntryScreenEvents.PinEntered, parentEventEmitter: (RegistrationFlowEvent) -> Unit): PinEntryState { Log.d(TAG, "[PinEntered] Attempting to restore master key from SVR...") - val restoreResult = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = true) + val restoreResult = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = true, isRegistered = false) val masterKey: MasterKey = when (restoreResult) { is RequestResult.Success -> { diff --git a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModel.kt b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModel.kt index cec930767f..8429ea23f5 100644 --- a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModel.kt +++ b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModel.kt @@ -117,7 +117,7 @@ class PinEntryForSmsBypassViewModel( return state } - return when (val result = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = false)) { + return when (val result = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = false, isRegistered = false)) { is RequestResult.Success -> { Log.i(TAG, "[PinEntered] Successfully restored master key from SVR.") parentEventEmitter(RegistrationFlowEvent.MasterKeyRestoredFromSvr(result.result.masterKey)) diff --git a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModel.kt b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModel.kt index b040798120..b7c59247e8 100644 --- a/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModel.kt +++ b/feature/registration/src/main/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModel.kt @@ -139,7 +139,7 @@ class PinEntryForSvrRestoreViewModel( } } - return when (val result = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = false)) { + return when (val result = repository.restoreMasterKeyFromSvr(svrCredentials, event.pin, forRegistrationLock = false, isRegistered = true)) { is RequestResult.Success -> { Log.i(TAG, "[PinEntered] Successfully restored master key from SVR.") repository.enqueueSvrResetGuessCountJob() diff --git a/feature/registration/src/test/java/org/signal/registration/RegistrationRepositoryTest.kt b/feature/registration/src/test/java/org/signal/registration/RegistrationRepositoryTest.kt index 6c1ec351cf..74de6fce70 100644 --- a/feature/registration/src/test/java/org/signal/registration/RegistrationRepositoryTest.kt +++ b/feature/registration/src/test/java/org/signal/registration/RegistrationRepositoryTest.kt @@ -7,8 +7,12 @@ package org.signal.registration import android.content.Context import assertk.assertThat +import assertk.assertions.isEmpty import assertk.assertions.isEqualTo import assertk.assertions.isInstanceOf +import assertk.assertions.isNotNull +import assertk.assertions.isNull +import assertk.assertions.isTrue import io.mockk.mockk import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.test.runTest @@ -17,9 +21,12 @@ import org.junit.Test import org.signal.core.models.AccountEntropyPool import org.signal.core.util.logging.Log import org.signal.libsignal.net.RequestResult +import org.signal.network.api.RegistrationApiV2.SvrCredentials import org.signal.registration.NetworkController.GetBackupInfoError import org.signal.registration.NetworkController.GetBackupInfoResponse +import org.signal.registration.NetworkController.MasterKeyResponse import org.signal.registration.NetworkController.ReserveBackupIdError +import org.signal.registration.NetworkController.RestoreMasterKeyError import org.signal.registration.fakes.FakeNetworkController import org.signal.registration.fakes.FakeStorageController import org.signal.registration.fakes.SystemOutLogger @@ -33,19 +40,23 @@ import kotlin.time.Duration.Companion.seconds class RegistrationRepositoryTest { private lateinit var networkController: FakeNetworkController + private lateinit var storageController: FakeStorageController private lateinit var repository: RegistrationRepository private val aep = AccountEntropyPool.generate() + private val masterKey = aep.deriveMasterKey() + private val svrCredentials = SvrCredentials(username = "user", password = "pass") private val backupInfo = GetBackupInfoResponse(cdn = 3, backupDir = "dir", mediaDir = "media", backupName = "backup", usedSpace = 1024L) @Before fun setup() { Log.initialize(SystemOutLogger()) networkController = FakeNetworkController() + storageController = FakeStorageController() repository = RegistrationRepository( context = mockk(relaxed = true), networkController = networkController, - storageController = FakeStorageController(), + storageController = storageController, isLinkAndSyncAvailable = false ) } @@ -132,4 +143,51 @@ class RegistrationRepositoryTest { assertThat((result as RequestResult.ApplicationError).cause).isEqualTo(cause) } + + // ==================== restoreMasterKeyFromSvr ==================== + + @Test + fun `restoreMasterKeyFromSvr commits the restored data when already registered`() = runTest { + networkController.onRestoreMasterKeyFromSvr = { RequestResult.Success(MasterKeyResponse(masterKey)) } + + val result = repository.restoreMasterKeyFromSvr(svrCredentials, pin = "1234", forRegistrationLock = false, isRegistered = true) + + assertThat(result).isInstanceOf(RequestResult.Success::class) + assertThat(storageController.committedData).isNotNull() + assertThat(storageController.committedData!!.pin).isEqualTo("1234") + } + + @Test + fun `restoreMasterKeyFromSvr does not commit the restored data when not yet registered`() = runTest { + networkController.onRestoreMasterKeyFromSvr = { RequestResult.Success(MasterKeyResponse(masterKey)) } + + val result = repository.restoreMasterKeyFromSvr(svrCredentials, pin = "1234", forRegistrationLock = false, isRegistered = false) + + assertThat(result).isInstanceOf(RequestResult.Success::class) + assertThat(storageController.committedData).isNull() + } + + @Test + fun `restoreMasterKeyFromSvr still records the in-progress data when not yet registered`() = runTest { + networkController.onRestoreMasterKeyFromSvr = { RequestResult.Success(MasterKeyResponse(masterKey)) } + + repository.restoreMasterKeyFromSvr(svrCredentials, pin = "1234", forRegistrationLock = true, isRegistered = false) + + val inProgress = storageController.readInProgressRegistrationData() + assertThat(inProgress.pin).isEqualTo("1234") + assertThat(inProgress.masterKeyForInitialDataRestore?.toByteArray()?.toList()).isEqualTo(masterKey.serialize().toList()) + assertThat(inProgress.registrationLockEnabled).isTrue() + assertThat(inProgress.svrCredentials.map { it.username }).isEqualTo(listOf("user")) + } + + @Test + fun `restoreMasterKeyFromSvr does not record or commit anything on failure`() = runTest { + networkController.onRestoreMasterKeyFromSvr = { RequestResult.NonSuccess(RestoreMasterKeyError.WrongPin(triesRemaining = 5)) } + + val result = repository.restoreMasterKeyFromSvr(svrCredentials, pin = "1234", forRegistrationLock = false, isRegistered = true) + + assertThat(result).isInstanceOf(RequestResult.NonSuccess::class) + assertThat(storageController.committedData).isNull() + assertThat(storageController.readInProgressRegistrationData().pin).isEmpty() + } } diff --git a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModelTest.kt b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModelTest.kt index dea0592522..21d379d744 100644 --- a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModelTest.kt +++ b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForRegistrationLockViewModelTest.kt @@ -84,7 +84,7 @@ class PinEntryForRegistrationLockViewModelTest { val registerResponse = createRegisterAccountResponse() val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.Success(registerResponse to keyMaterial) @@ -106,7 +106,7 @@ class PinEntryForRegistrationLockViewModelTest { val registerResponse = createRegisterAccountResponse(reregistration = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.Success(registerResponse to keyMaterial) @@ -133,7 +133,7 @@ class PinEntryForRegistrationLockViewModelTest { parentState.value = parentState.value.copy(preExistingRegistrationData = mockk(relaxed = true)) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.Success(registerResponse to keyMaterial) @@ -155,7 +155,7 @@ class PinEntryForRegistrationLockViewModelTest { unverifiedRestoredAep = restoreAep ) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.Success(createRegisterAccountResponse() to keyMaterial) @@ -182,7 +182,7 @@ class PinEntryForRegistrationLockViewModelTest { unverifiedRestoredAep = restoreAep ) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.Success(createRegisterAccountResponse() to keyMaterial) @@ -202,7 +202,7 @@ class PinEntryForRegistrationLockViewModelTest { val triesRemaining = 3 val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(triesRemaining) ) @@ -218,7 +218,7 @@ class PinEntryForRegistrationLockViewModelTest { fun `PinEntered with wrong PIN and no tries remaining navigates to AccountLocked`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(0) ) @@ -238,7 +238,7 @@ class PinEntryForRegistrationLockViewModelTest { fun `PinEntered with no SVR data navigates to AccountLocked`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.NoDataFound ) @@ -259,7 +259,7 @@ class PinEntryForRegistrationLockViewModelTest { fun `PinEntered with network error when restoring master key returns NetworkError event`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.RetryableNetworkError(java.io.IOException("Network error")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -273,7 +273,7 @@ class PinEntryForRegistrationLockViewModelTest { fun `PinEntered with application error when restoring master key returns UnknownError event`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.ApplicationError(RuntimeException("Unexpected")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -293,7 +293,7 @@ class PinEntryForRegistrationLockViewModelTest { sessionE164 = null ) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -315,7 +315,7 @@ class PinEntryForRegistrationLockViewModelTest { sessionE164 = "+15551234567" ) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.Success(registerResponse to keyMaterial) @@ -343,7 +343,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -362,7 +362,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -383,7 +383,7 @@ class PinEntryForRegistrationLockViewModelTest { val registrationLockData = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -408,7 +408,7 @@ class PinEntryForRegistrationLockViewModelTest { val retryAfter = 30.seconds val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -428,7 +428,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -448,7 +448,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -468,7 +468,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.RetryableNetworkError(java.io.IOException("Network error")) @@ -486,7 +486,7 @@ class PinEntryForRegistrationLockViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns RequestResult.ApplicationError(RuntimeException("Unexpected")) @@ -562,4 +562,21 @@ class PinEntryForRegistrationLockViewModelTest { entitlements = null, reregistration = reregistration ) + + @Test + fun `PinEntered restores with isRegistered false because registration has not happened yet`() = runTest { + val masterKey = mockk(relaxed = true) + val keyMaterial = mockk(relaxed = true) + val registerResponse = createRegisterAccountResponse() + val initialState = PinEntryState(mode = PinEntryState.Mode.RegistrationLock) + + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), any(), any()) } returns + RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) + coEvery { mockRepository.registerAccountWithSession(any(), any(), any(), any()) } returns + RequestResult.Success(registerResponse to keyMaterial) + + viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) + + coVerify { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = true, isRegistered = false) } + } } diff --git a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModelTest.kt b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModelTest.kt index fc8fe06bff..512c125a6b 100644 --- a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModelTest.kt +++ b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSmsBypassViewModelTest.kt @@ -71,7 +71,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.Success(mockk(relaxed = true)) @@ -90,7 +90,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.Success(mockk(relaxed = true)) @@ -105,7 +105,7 @@ class PinEntryForSmsBypassViewModelTest { val triesRemaining = 3 val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(triesRemaining) ) @@ -121,7 +121,7 @@ class PinEntryForSmsBypassViewModelTest { fun `PinEntered with no SVR data emits RecoveryPasswordInvalid and navigates back`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.NoDataFound ) @@ -138,7 +138,7 @@ class PinEntryForSmsBypassViewModelTest { fun `PinEntered with network error restoring master key returns NetworkError event`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.RetryableNetworkError(IOException("Network error")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -152,7 +152,7 @@ class PinEntryForSmsBypassViewModelTest { fun `PinEntered with application error restoring master key returns UnknownError event`() = runTest { val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.ApplicationError(RuntimeException("Unexpected")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -180,7 +180,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.RetryableNetworkError(IOException("Network error")) @@ -198,7 +198,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.ApplicationError(RuntimeException("Unexpected")) @@ -216,7 +216,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -236,7 +236,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -257,7 +257,7 @@ class PinEntryForSmsBypassViewModelTest { val retryAfter = 30.seconds val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -278,7 +278,7 @@ class PinEntryForSmsBypassViewModelTest { val registrationLockData = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) // First call (without reglock) returns RegistrationLock error, second call (with reglock) succeeds coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), registrationLock = null, any()) } returns @@ -303,7 +303,7 @@ class PinEntryForSmsBypassViewModelTest { val registrationLockData = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) // Both calls (with and without reglock) return RegistrationLock error coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns @@ -325,7 +325,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -345,7 +345,7 @@ class PinEntryForSmsBypassViewModelTest { val masterKey = mockk(relaxed = true) val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns RequestResult.NonSuccess( @@ -400,4 +400,19 @@ class PinEntryForSmsBypassViewModelTest { assertThat(emittedStates.last().isAlphanumericKeyboard).isEqualTo(false) } + + @Test + fun `PinEntered restores with isRegistered false because registration has not happened yet`() = runTest { + val masterKey = mockk(relaxed = true) + val initialState = PinEntryState(mode = PinEntryState.Mode.SmsBypass, e164 = "+15551234567") + + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), any(), any()) } returns + RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) + coEvery { mockRepository.registerAccountWithRecoveryPassword(any(), any(), any(), any()) } returns + RequestResult.Success(mockk(relaxed = true)) + + viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) + + coVerify { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = false) } + } } diff --git a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModelTest.kt b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModelTest.kt index 29f4979358..65666fe4f0 100644 --- a/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModelTest.kt +++ b/feature/registration/src/test/java/org/signal/registration/screens/pinentry/PinEntryForSvrRestoreViewModelTest.kt @@ -73,7 +73,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -161,7 +161,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(triesRemaining) ) @@ -192,7 +192,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(3) ) @@ -212,7 +212,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.WrongPin(3) ) @@ -232,7 +232,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.NonSuccess( NetworkController.RestoreMasterKeyError.NoDataFound ) @@ -280,7 +280,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.RetryableNetworkError(java.io.IOException("Network error")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -300,7 +300,7 @@ class PinEntryForSvrRestoreViewModelTest { coEvery { mockRepository.getSvrCredentials() } returns RequestResult.Success(svrCredentials) - coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false) } returns + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } returns RequestResult.ApplicationError(RuntimeException("Unexpected")) viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) @@ -360,4 +360,23 @@ class PinEntryForSvrRestoreViewModelTest { requestedInformation = requestedInformation, verified = verified ) + + @Test + fun `PinEntered restores with isRegistered true because this screen is only reached post-registration`() = runTest { + val masterKey = mockk(relaxed = true) + val svrCredentials = SvrCredentials( + username = "test-username", + password = "test-password" + ) + val initialState = PinEntryState(mode = PinEntryState.Mode.SvrRestore) + + coEvery { mockRepository.getSvrCredentials() } returns + RequestResult.Success(svrCredentials) + coEvery { mockRepository.restoreMasterKeyFromSvr(any(), any(), any(), any()) } returns + RequestResult.Success(NetworkController.MasterKeyResponse(masterKey)) + + viewModel.applyEvent(initialState, PinEntryScreenEvents.PinEntered("123456"), parentEventEmitter, stateEmitter) + + coVerify { mockRepository.restoreMasterKeyFromSvr(any(), any(), forRegistrationLock = false, isRegistered = true) } + } }