feat(auth, device-migration)!: Perform post-login steps after obtaining forked session.

MIGRATION: AccountDatabase.MIGRATION_10
This commit is contained in:
Mateusz Armatys
2025-04-10 18:01:55 +02:00
parent 606b9b678c
commit 080dc333af
64 changed files with 7984 additions and 275 deletions
@@ -203,7 +203,7 @@ abstract class AccountManagerDatabase :
companion object {
const val name = "db-account-manager"
const val version = 60
const val version = 61
val migrations = listOf(
AccountManagerDatabaseMigrations.MIGRATION_1_2,
@@ -265,6 +265,7 @@ abstract class AccountManagerDatabase :
AccountManagerDatabaseMigrations.MIGRATION_57_58,
AccountManagerDatabaseMigrations.MIGRATION_58_59,
AccountManagerDatabaseMigrations.MIGRATION_59_60,
AccountManagerDatabaseMigrations.MIGRATION_60_61,
)
fun databaseBuilder(context: Context): Builder<AccountManagerDatabase> =
@@ -417,4 +417,10 @@ object AccountManagerDatabaseMigrations {
UserSettingsDatabase.MIGRATION_8.migrate(db)
}
}
val MIGRATION_60_61 = object : Migration(60, 61) {
override fun migrate(db: SupportSQLiteDatabase) {
AccountDatabase.MIGRATION_10.migrate(db)
}
}
}
@@ -187,5 +187,31 @@ interface AccountDatabase : Database {
}
}
/**
* Add [me.proton.core.account.data.entity.SessionDetailsEntity.passphrase] column.
* Column [me.proton.core.account.data.entity.SessionDetailsEntity.initialEventId] is now nullable.
*/
val MIGRATION_10 = object : DatabaseMigration {
override fun migrate(database: SupportSQLiteDatabase) {
val tableName = "SessionDetailsEntity"
database.addTableColumn(
table = tableName,
column = "passphrase",
type = "BLOB"
)
// `initialEventId` is now nullable.
database.recreateTable(
table = tableName,
createTable = {
execSQL("CREATE TABLE IF NOT EXISTS `${tableName}` (`sessionId` TEXT NOT NULL, `initialEventId` TEXT, `requiredAccountType` TEXT NOT NULL, `secondFactorEnabled` INTEGER NOT NULL, `twoPassModeEnabled` INTEGER NOT NULL, `passphrase` BLOB, `password` TEXT, `fido2AuthenticationOptionsJson` TEXT, PRIMARY KEY(`sessionId`), FOREIGN KEY(`sessionId`) REFERENCES `SessionEntity`(`sessionId`) ON UPDATE NO ACTION ON DELETE CASCADE )")
},
createIndices = {
execSQL("CREATE INDEX IF NOT EXISTS `index_SessionDetailsEntity_sessionId` ON `${tableName}` (`sessionId`)")
}
)
}
}
}
}
@@ -37,6 +37,6 @@ abstract class SessionDetailsDao : BaseDao<SessionDetailsEntity>() {
@Query("DELETE FROM SessionDetailsEntity WHERE sessionId = :sessionId")
abstract suspend fun delete(sessionId: SessionId)
@Query("UPDATE SessionDetailsEntity SET password = null WHERE sessionId = :sessionId")
abstract suspend fun clearPassword(sessionId: SessionId)
@Query("UPDATE SessionDetailsEntity SET passphrase = null, password = null WHERE sessionId = :sessionId")
abstract suspend fun clearAuthSecrets(sessionId: SessionId)
}
@@ -23,6 +23,7 @@ import androidx.room.ForeignKey
import androidx.room.Index
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.account.domain.entity.SessionDetails
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.network.domain.session.SessionId
import me.proton.core.util.kotlin.deserializeOrNull
@@ -41,10 +42,11 @@ import me.proton.core.util.kotlin.deserializeOrNull
)
data class SessionDetailsEntity(
val sessionId: SessionId,
val initialEventId: String,
val initialEventId: String?,
val requiredAccountType: AccountType,
val secondFactorEnabled: Boolean,
val twoPassModeEnabled: Boolean,
val passphrase: EncryptedByteArray?,
val password: EncryptedString?,
val fido2AuthenticationOptionsJson: String? = null
) {
@@ -53,6 +55,7 @@ data class SessionDetailsEntity(
requiredAccountType = requiredAccountType,
secondFactorEnabled = secondFactorEnabled,
twoPassModeEnabled = twoPassModeEnabled,
passphrase = passphrase,
password = password,
fido2AuthenticationOptionsJson = fido2AuthenticationOptionsJson
)
@@ -278,13 +278,14 @@ class AccountRepositoryImpl @Inject constructor(
requiredAccountType = details.requiredAccountType,
secondFactorEnabled = details.secondFactorEnabled,
twoPassModeEnabled = details.twoPassModeEnabled,
passphrase = details.passphrase,
password = details.password,
fido2AuthenticationOptionsJson = details.fido2AuthenticationOptionsJson
)
)
override suspend fun clearSessionDetails(sessionId: SessionId) =
sessionDetailsDao.clearPassword(sessionId = sessionId)
sessionDetailsDao.clearAuthSecrets(sessionId = sessionId)
override suspend fun addMigration(userId: UserId, migration: String) =
db.inTransaction {
@@ -260,7 +260,7 @@ class AccountRepositoryImplTest {
fun `clear account session details`() = runTest {
accountRepository.clearSessionDetails(testAccountEntity.toAccount(ad).sessionId!!)
coVerify(exactly = 1) { sessionDetailsDao.clearPassword(any()) }
coVerify(exactly = 1) { sessionDetailsDao.clearAuthSecrets(any()) }
}
@Test
@@ -278,6 +278,7 @@ class AccountRepositoryImplTest {
requiredAccountType = AccountType.Internal,
secondFactorEnabled = true,
twoPassModeEnabled = true,
passphrase = null,
password = "encrypted-password",
fido2AuthenticationOptionsJson = null
)
@@ -293,6 +294,7 @@ class AccountRepositoryImplTest {
requiredAccountType = AccountType.Internal,
secondFactorEnabled = true,
twoPassModeEnabled = true,
passphrase = null,
password = "encrypted-password",
fido2AuthenticationOptionsJson = null
)
@@ -500,6 +502,7 @@ class AccountRepositoryImplTest {
requiredAccountType = AccountType.Internal,
secondFactorEnabled = true,
twoPassModeEnabled = true,
passphrase = null,
password = "encrypted-password",
fido2AuthenticationOptionsJson = null
)
@@ -513,6 +516,7 @@ class AccountRepositoryImplTest {
requiredAccountType = AccountType.Internal,
secondFactorEnabled = true,
twoPassModeEnabled = true,
passphrase = null,
password = "encrypted-password",
fido2AuthenticationOptionsJson = null
)
@@ -18,6 +18,7 @@
package me.proton.core.account.domain.entity
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.domain.entity.UserId
import me.proton.core.network.domain.session.SessionId
@@ -43,10 +44,11 @@ data class AccountMetadataDetails(
)
data class SessionDetails(
val initialEventId: String,
val initialEventId: String?,
val requiredAccountType: AccountType,
val secondFactorEnabled: Boolean,
val twoPassModeEnabled: Boolean,
val passphrase: EncryptedByteArray?,
val password: EncryptedString?,
val fido2AuthenticationOptionsJson: String?
)
@@ -0,0 +1,32 @@
/*
* Copyright (c) 2025 Proton AG
* This file is part of Proton AG and ProtonCore.
*
* ProtonCore is free software: you can redistribute it and/or modify
* it under the terms of the GNU General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* ProtonCore is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU General Public License for more details.
*
* You should have received a copy of the GNU General Public License
* along with ProtonCore. If not, see <https://www.gnu.org/licenses/>.
*/
package me.proton.core.auth.domain.entity
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
sealed interface EncryptedAuthSecret {
data object Absent : EncryptedAuthSecret
@JvmInline
value class Password(val password: EncryptedString) : EncryptedAuthSecret
@JvmInline
value class Passphrase(val passphrase: EncryptedByteArray) : EncryptedAuthSecret
}
@@ -25,7 +25,7 @@ import kotlinx.serialization.Serializable
@Serializable
data class SessionForkPayloadWithKey(
@SerialName("keyPassword")
val keyPassword: String,
val keyPassword: String?,
@EncodeDefault
@SerialName("type")
@@ -67,6 +67,7 @@ class CreateLoginLessSession @Inject constructor(
requiredAccountType = requiredAccountType,
secondFactorEnabled = sessionInfo.isSecondFactorNeeded,
twoPassModeEnabled = sessionInfo.isTwoPassModeNeeded,
passphrase = null,
password = null,
fido2AuthenticationOptionsJson = sessionInfo.getFido2AuthOptions()?.serialize()
)
@@ -74,6 +74,7 @@ class CreateLoginSession @Inject constructor(
requiredAccountType = requiredAccountType,
secondFactorEnabled = sessionInfo.isSecondFactorNeeded,
twoPassModeEnabled = sessionInfo.isTwoPassModeNeeded,
passphrase = null,
password = password,
fido2AuthenticationOptionsJson = sessionInfo.getFido2AuthOptions()?.serialize()
)
@@ -0,0 +1,71 @@
/*
* Copyright (c) 2025 Proton AG
* This file is part of Proton AG and ProtonCore.
*
* ProtonCore is free software: you can redistribute it and/or modify
* it under the terms of the GNU General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* ProtonCore is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU General Public License for more details.
*
* You should have received a copy of the GNU General Public License
* along with ProtonCore. If not, see <https://www.gnu.org/licenses/>.
*/
package me.proton.core.auth.domain.usecase
import me.proton.core.account.domain.entity.Account
import me.proton.core.account.domain.entity.AccountDetails
import me.proton.core.account.domain.entity.AccountState
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.account.domain.entity.SessionDetails
import me.proton.core.account.domain.entity.SessionState
import me.proton.core.account.domain.repository.AccountRepository
import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.network.domain.session.Session
import me.proton.core.user.domain.UserManager
import javax.inject.Inject
class CreateLoginSessionFromFork @Inject constructor(
private val accountRepository: AccountRepository,
private val accountWorkflowHandler: AccountWorkflowHandler,
private val userManager: UserManager,
) {
suspend operator fun invoke(
accountType: AccountType,
passphrase: EncryptedByteArray?,
session: Session.Authenticated
) {
val account = Account(
userId = session.userId,
username = null,
email = null,
state = AccountState.NotReady,
sessionId = session.sessionId,
sessionState = SessionState.Authenticated,
details = AccountDetails(
session = SessionDetails(
initialEventId = null,
requiredAccountType = accountType,
secondFactorEnabled = false,
twoPassModeEnabled = false,
passphrase = passphrase,
password = null,
fido2AuthenticationOptionsJson = null,
)
)
)
accountWorkflowHandler.handleSession(account, session)
val user = userManager.getUser(session.userId)
accountRepository.createOrUpdateAccountSession(
account = account.copy(username = user.name, email = user.email),
session = session
)
}
}
@@ -62,6 +62,7 @@ class CreateLoginSsoSession @Inject constructor(
requiredAccountType = requiredAccountType,
secondFactorEnabled = sessionInfo.isSecondFactorNeeded,
twoPassModeEnabled = sessionInfo.isTwoPassModeNeeded,
passphrase = null,
password = null,
fido2AuthenticationOptionsJson = sessionInfo.getFido2AuthOptions()?.serialize()
)
@@ -23,6 +23,7 @@ import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.accountmanager.domain.SessionManager
import me.proton.core.auth.domain.LogTag
import me.proton.core.auth.domain.entity.BillingDetails
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.domain.entity.UserId
import me.proton.core.payment.domain.MAX_PLAN_QUANTITY
@@ -95,6 +96,30 @@ class PostLoginAccountSetup @Inject constructor(
onSetupSuccess: (suspend () -> Unit)? = null,
billingDetails: BillingDetails? = null,
internalAddressDomain: String? = null
): Result {
return invoke(
userId = userId,
encryptedAuthSecret = EncryptedAuthSecret.Password(encryptedPassword),
requiredAccountType = requiredAccountType,
isSecondFactorNeeded = isSecondFactorNeeded,
isTwoPassModeNeeded = isTwoPassModeNeeded,
temporaryPassword = temporaryPassword,
onSetupSuccess = onSetupSuccess,
billingDetails = billingDetails,
internalAddressDomain = internalAddressDomain,
)
}
suspend operator fun invoke(
userId: UserId,
encryptedAuthSecret: EncryptedAuthSecret,
requiredAccountType: AccountType,
isSecondFactorNeeded: Boolean,
isTwoPassModeNeeded: Boolean,
temporaryPassword: Boolean,
onSetupSuccess: (suspend () -> Unit)? = null,
billingDetails: BillingDetails? = null,
internalAddressDomain: String? = null
): Result {
// Flows not using PurchaseStateHandler pass billingDetails.
subscribeAnyPendingBilling(billingDetails, userId)
@@ -120,22 +145,27 @@ class PostLoginAccountSetup @Inject constructor(
Result.Need.ChooseUsername(userId)
}
is SetupAccountCheck.Result.SetupPrimaryKeysNeeded -> {
setupPrimaryKeys.invoke(
userId,
encryptedPassword,
requiredAccountType,
internalAddressDomain
)
unlockUserPrimaryKey(
userId,
encryptedPassword,
onSetupSuccess
)
val encryptedPassword = encryptedAuthSecret as? EncryptedAuthSecret.Password
if (encryptedPassword != null) {
setupPrimaryKeys.invoke(
userId,
encryptedPassword.password,
requiredAccountType,
internalAddressDomain
)
unlockUserPrimaryKey(
userId,
encryptedAuthSecret,
onSetupSuccess
)
} else {
Result.Error.UnlockPrimaryKeyError(UserManager.UnlockResult.Error.PrimaryKeyInvalidPassphrase)
}
}
is SetupAccountCheck.Result.SetupExternalAddressKeysNeeded -> {
unlockUserPrimaryKey(
userId,
encryptedPassword,
encryptedAuthSecret,
onSetupSuccess
) {
setupExternalAddressKeys.invoke(userId)
@@ -144,7 +174,7 @@ class PostLoginAccountSetup @Inject constructor(
is SetupAccountCheck.Result.SetupInternalAddressNeeded -> {
unlockUserPrimaryKey(
userId,
encryptedPassword,
encryptedAuthSecret,
onSetupSuccess
) {
setupInternalAddress.invoke(userId, internalAddressDomain)
@@ -153,7 +183,7 @@ class PostLoginAccountSetup @Inject constructor(
is SetupAccountCheck.Result.NoSetupNeeded -> {
unlockUserPrimaryKey(
userId,
encryptedPassword,
encryptedAuthSecret,
onSetupSuccess
)
}
@@ -207,11 +237,11 @@ class PostLoginAccountSetup @Inject constructor(
private suspend fun unlockUserPrimaryKey(
userId: UserId,
password: EncryptedString,
secret: EncryptedAuthSecret,
onSetupSuccess: (suspend () -> Unit)?,
onUnlockSuccess: (suspend () -> Unit)? = null,
): Result {
return when (val result = unlockUserPrimaryKey.invoke(userId, password)) {
return when (val result = unlockUserPrimaryKey.invoke(userId, secret)) {
is UserManager.UnlockResult.Success -> {
// Invoke unlock success action.
onUnlockSuccess?.invoke()
@@ -18,13 +18,15 @@
package me.proton.core.auth.domain.usecase
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.account.domain.repository.AccountRepository
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.crypto.common.keystore.KeyStoreCrypto
import me.proton.core.crypto.common.keystore.decrypt
import me.proton.core.crypto.common.keystore.use
import me.proton.core.domain.entity.Product
import me.proton.core.domain.entity.UserId
import me.proton.core.user.domain.UserManager
import me.proton.core.user.domain.entity.Type
import me.proton.core.user.domain.entity.UserKey
import me.proton.core.user.domain.extension.hasKeys
import me.proton.core.util.kotlin.coroutine.result
@@ -38,6 +40,7 @@ import javax.inject.Inject
* - For VPN: this function always return UnlockResult.Success.
*/
class UnlockUserPrimaryKey @Inject constructor(
private val accountRepository: AccountRepository,
private val userManager: UserManager,
private val keyStoreCrypto: KeyStoreCrypto,
private val product: Product
@@ -54,16 +57,26 @@ class UnlockUserPrimaryKey @Inject constructor(
)
suspend operator fun invoke(
userId: UserId,
password: EncryptedString
secret: EncryptedAuthSecret
): UserManager.UnlockResult {
return when {
product == Product.Vpn -> UserManager.UnlockResult.Success
shouldSkipForVpn(userId) -> UserManager.UnlockResult.Success
!userManager.getUser(userId).hasKeys() -> UserManager.UnlockResult.Success
else -> password.decrypt(keyStoreCrypto).toByteArray().use {
userManager.unlockWithPassword(userId, it)
else -> when (secret) {
is EncryptedAuthSecret.Absent -> UserManager.UnlockResult.Error.PrimaryKeyInvalidPassphrase
is EncryptedAuthSecret.Passphrase -> userManager.unlockWithPassphrase(userId, secret.passphrase)
is EncryptedAuthSecret.Password -> secret.password.decrypt(keyStoreCrypto).toByteArray().use {
userManager.unlockWithPassword(userId, it)
}
}
}.let { unlockResult ->
result("unlockUserPrimaryKey") { unlockResult }
}
}
private suspend fun shouldSkipForVpn(userId: UserId): Boolean {
if (product != Product.Vpn) return false
val account = accountRepository.getAccountOrNull(userId)
return account?.details?.session?.twoPassModeEnabled == true || !userManager.getUser(userId).hasKeys()
}
}
@@ -31,6 +31,7 @@ import me.proton.core.crypto.common.keystore.decrypt
import me.proton.core.crypto.common.keystore.encrypt
import me.proton.core.util.kotlin.DispatcherProvider
import me.proton.core.util.kotlin.deserialize
import me.proton.core.util.kotlin.takeIfNotBlank
import javax.inject.Inject
class DecryptPassphrasePayload @Inject constructor(
@@ -46,7 +47,7 @@ class DecryptPassphrasePayload @Inject constructor(
encryptionKey: EncryptedByteArray,
aesCipherGCMTagBits: Int,
aesCipherIvBytes: Int
): EncryptedByteArray = withContext(dispatcherProvider.Comp) {
): EncryptedByteArray? = withContext(dispatcherProvider.Comp) {
val aeadCrypto = cryptoContext.aeadCryptoFactory.create(
keyAlgorithm = DEFAULT_AES_KEY_ALGORITHM,
transformation = DEFAULT_AES_GCM_CIPHER_TRANSFORMATION,
@@ -57,7 +58,7 @@ class DecryptPassphrasePayload @Inject constructor(
payload.decrypt(aeadCrypto, key = it.array)
}
val payloadWithKey = decryptedPayload.deserialize<SessionForkPayloadWithKey>()
val key = PlainByteArray(payloadWithKey.keyPassword.toByteArray())
key.encrypt(cryptoContext.keyStoreCrypto)
val key = payloadWithKey.keyPassword?.takeIfNotBlank()?.let { PlainByteArray(it.toByteArray()) }
key?.encrypt(cryptoContext.keyStoreCrypto)
}
}
@@ -54,15 +54,15 @@ class GetEncryptedPassphrasePayload @Inject constructor(
authTagBits = aesCipherGCMTagBits,
ivBytes = aesCipherIvBytes
)
requireNotNull(passphraseRepository.getPassphrase(userId))
.decrypt(cryptoContext.keyStoreCrypto)
.use { decryptedPassphrase ->
val serializedPayload = SessionForkPayloadWithKey(
keyPassword = String(decryptedPassphrase.array)
).serialize()
encryptionKey.decrypt(cryptoContext.keyStoreCrypto).use { key ->
serializedPayload.encrypt(aeadCrypto, key = key.array)
}
// Note: passphrase may be `null` on VPN.
passphraseRepository.getPassphrase(userId)?.decrypt(cryptoContext.keyStoreCrypto).use { decryptedPassphrase ->
val serializedPayload = SessionForkPayloadWithKey(
keyPassword = decryptedPassphrase?.let { String(it.array) }
).serialize()
encryptionKey.decrypt(cryptoContext.keyStoreCrypto).use { key ->
serializedPayload.encrypt(aeadCrypto, key = key.array)
}
}
}
}
@@ -26,6 +26,7 @@ import io.mockk.mockk
import kotlinx.coroutines.test.runTest
import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.accountmanager.domain.SessionManager
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.entity.SessionInfo
import me.proton.core.auth.domain.entity.UnprivatizationInfo
import me.proton.core.auth.domain.repository.AuthDeviceRepository
@@ -92,7 +93,7 @@ class PostLoginSsoAccountSetupTest {
coEvery { this@mockk.invoke(any()) } returns PostLoginAccountSetup.UserCheckResult.Success
}
unlockUserPrimaryKey = mockk {
coEvery { this@mockk.invoke(any(), any<EncryptedString>()) } returns UserManager.UnlockResult.Success
coEvery { this@mockk.invoke(any(), any<EncryptedAuthSecret>()) } returns UserManager.UnlockResult.Success
}
userManager = mockk {
coEvery { this@mockk.getUser(any(), any()) } returns user
@@ -20,9 +20,13 @@ package me.proton.core.auth.domain.usecase
import io.mockk.MockKAnnotations
import io.mockk.coEvery
import io.mockk.coVerify
import io.mockk.every
import io.mockk.impl.annotations.MockK
import io.mockk.mockk
import me.proton.core.account.domain.repository.AccountRepository
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.crypto.common.keystore.KeyStoreCrypto
import me.proton.core.domain.entity.Product
@@ -34,6 +38,9 @@ import kotlin.test.Test
import kotlin.test.assertEquals
internal class UnlockUserPrimaryKeyTest {
@MockK
private lateinit var accountRepository: AccountRepository
@MockK
private lateinit var keyStoreCrypto: KeyStoreCrypto
@@ -46,13 +53,72 @@ internal class UnlockUserPrimaryKeyTest {
}
@Test
fun vpnSuccess_result() = runTestWithResultContext {
fun vpnSuccess_result_2pass() = runTestWithResultContext {
// GIVEN
coEvery { accountRepository.getAccountOrNull(any<UserId>()) } returns mockk {
every { details } returns mockk {
every { session } returns mockk {
every { twoPassModeEnabled } returns true
}
}
}
// WHEN
runTested(Product.Vpn)
// THEN
val result = assertSingleResult("unlockUserPrimaryKey")
assertEquals(UserManager.UnlockResult.Success, result.getOrThrow())
coVerify(exactly = 0) { userManager.unlockWithPassphrase(any(), any()) }
coVerify(exactly = 0) { userManager.unlockWithPassword(any(), any()) }
}
@Test
fun vpnSuccess_result_noKeys() = runTestWithResultContext {
// GIVEN
coEvery { accountRepository.getAccountOrNull(any<UserId>()) } returns mockk {
every { details } returns mockk {
every { session } returns mockk {
every { twoPassModeEnabled } returns false
}
}
}
coEvery { userManager.getUser(any(), any()) } returns mockk {
every { keys } returns emptyList()
}
// WHEN
runTested(Product.Vpn)
// THEN
val result = assertSingleResult("unlockUserPrimaryKey")
assertEquals(UserManager.UnlockResult.Success, result.getOrThrow())
coVerify(exactly = 0) { userManager.unlockWithPassphrase(any(), any()) }
coVerify(exactly = 0) { userManager.unlockWithPassword(any(), any()) }
}
@Test
fun vpnSuccess_result_unlocks() = runTestWithResultContext {
// GIVEN
coEvery { accountRepository.getAccountOrNull(any<UserId>()) } returns mockk {
every { details } returns mockk {
every { session } returns mockk {
every { twoPassModeEnabled } returns false
}
}
}
coEvery { userManager.getUser(any(), any()) } returns mockk {
every { keys } returns listOf(mockk())
}
coEvery { userManager.unlockWithPassphrase(any(), any<EncryptedByteArray>()) } returns
UserManager.UnlockResult.Success
// WHEN
runTested(Product.Vpn, authSecret = EncryptedAuthSecret.Passphrase(EncryptedByteArray(byteArrayOf(1, 2, 3))))
// THEN
val result = assertSingleResult("unlockUserPrimaryKey")
assertEquals(UserManager.UnlockResult.Success, result.getOrThrow())
coVerify { userManager.unlockWithPassphrase(any(), any()) }
}
@Test
@@ -88,8 +154,11 @@ internal class UnlockUserPrimaryKeyTest {
assertEquals(UserManager.UnlockResult.Error.NoPrimaryKey, result.getOrThrow())
}
private suspend fun runTested(product: Product = Product.Mail) {
val tested = UnlockUserPrimaryKey(userManager, keyStoreCrypto, product)
tested(UserId("test_user_id"), "test_password")
private suspend fun runTested(
product: Product = Product.Mail,
authSecret: EncryptedAuthSecret = EncryptedAuthSecret.Password("test_password")
) {
val tested = UnlockUserPrimaryKey(accountRepository, userManager, keyStoreCrypto, product)
tested(UserId("test_user_id"), authSecret)
}
}
+1
View File
@@ -55,6 +55,7 @@ dependencies {
project(Module.challengeDomain),
project(Module.challengePresentation),
project(Module.countryDomain),
project(Module.cryptoAndroid),
project(Module.cryptoCommon),
project(Module.deviceMigrationPresentation),
project(Module.domain),
@@ -25,10 +25,12 @@ import androidx.activity.result.ActivityResultLauncher
import dagger.hilt.android.qualifiers.ApplicationContext
import me.proton.core.account.domain.entity.Account
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.account.domain.entity.SessionDetails
import me.proton.core.auth.domain.feature.IsLoginTwoStepEnabled
import me.proton.core.auth.presentation.alert.confirmpass.StartConfirmPassword
import me.proton.core.auth.presentation.entity.AddAccountInput
import me.proton.core.auth.presentation.entity.AddAccountResult
import me.proton.core.auth.presentation.entity.ChooseAddressAuthSecret
import me.proton.core.auth.presentation.entity.ChooseAddressInput
import me.proton.core.auth.presentation.entity.ChooseAddressResult
import me.proton.core.auth.presentation.entity.DeviceSecretResult
@@ -53,6 +55,7 @@ import me.proton.core.auth.presentation.ui.StartLoginTwoStep
import me.proton.core.auth.presentation.ui.StartSecondFactor
import me.proton.core.auth.presentation.ui.StartSignup
import me.proton.core.auth.presentation.ui.StartTwoPassMode
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.domain.entity.UserId
import me.proton.core.network.domain.scopes.MissingScopeState
@@ -228,14 +231,14 @@ class AuthOrchestrator @Inject constructor(
private fun startChooseAddressWorkflow(
userId: UserId,
password: EncryptedString,
authSecret: ChooseAddressAuthSecret,
externalEmail: String,
isTwoPassModeNeeded: Boolean
) {
checkRegistered(chooseAddressLauncher).launch(
ChooseAddressInput(
userId.id,
password = password,
authSecret = authSecret,
recoveryEmail = externalEmail,
isTwoPassModeNeeded = isTwoPassModeNeeded
)
@@ -382,20 +385,28 @@ class AuthOrchestrator @Inject constructor(
val email = checkNotNull(account.email) {
"Email is null for startChooseAddressWorkflow."
}
val password = checkNotNull(account.details.session?.password) {
"Password is null for startChooseAddressWorkflow."
}
val authSecret = getChoosePasswordAuthSecret(account.details.session)
val twoPassModeEnabled = checkNotNull(account.details.session?.twoPassModeEnabled) {
"TwoPassModeEnabled is null for startChooseAddressWorkflow."
}
startChooseAddressWorkflow(
userId = account.userId,
password = password,
authSecret = authSecret,
externalEmail = email,
isTwoPassModeNeeded = twoPassModeEnabled
)
}
private fun getChoosePasswordAuthSecret(sessionDetails: SessionDetails?): ChooseAddressAuthSecret {
val passphrase = sessionDetails?.passphrase
val password = sessionDetails?.password
return when {
passphrase != null && password == null -> ChooseAddressAuthSecret.Passphrase(passphrase)
passphrase == null && password != null -> ChooseAddressAuthSecret.Password(password)
else -> error("Either passphrase or password must be set.")
}
}
/**
* Start the Device Secret workflow.
*
@@ -24,6 +24,8 @@ import kotlinx.parcelize.Parcelize
@Parcelize
sealed interface AuthHelpResult : Parcelable {
/** User has signed in using Easy Device Migration (QR code). */
@Parcelize
data class SignedInWithEdm(val userId: String) : AuthHelpResult
/** User has signed in using Easy Device Migration (QR code) and needs to change password. */
data object PasswordChangeNeededAfterEdm : AuthHelpResult
}
@@ -20,12 +20,22 @@ package me.proton.core.auth.presentation.entity
import android.os.Parcelable
import kotlinx.parcelize.Parcelize
import kotlinx.parcelize.TypeParceler
import me.proton.core.crypto.android.keystore.EncryptedByteArrayParceler
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
@Parcelize
data class ChooseAddressInput constructor(
val userId: String,
val password: EncryptedString,
val authSecret: ChooseAddressAuthSecret,
val recoveryEmail: String,
val isTwoPassModeNeeded: Boolean
) : Parcelable
@Parcelize
sealed class ChooseAddressAuthSecret : Parcelable {
@TypeParceler<EncryptedByteArray, EncryptedByteArrayParceler>
data class Passphrase(val passphrase: EncryptedByteArray) : ChooseAddressAuthSecret()
data class Password(val password: EncryptedString) : ChooseAddressAuthSecret()
}
@@ -53,17 +53,8 @@ class AuthHelpActivity : AuthActivity<ActivityAuthHelpBinding>(ActivityAuthHelpB
override fun onCreate(savedInstanceState: Bundle?) {
super.onCreate(savedInstanceState)
targetDeviceMigrationLauncher = registerForActivityResult(StartMigrationFromTargetDevice()) { result ->
if (result is TargetDeviceMigrationResult.NavigateToSignIn) {
setResult(RESULT_CANCELED)
finish()
} else if (result is TargetDeviceMigrationResult.SignedIn) {
setResult(RESULT_OK, Intent().apply {
putExtra(ARG_RESULT, AuthHelpResult.SignedInWithEdm(result.userId))
})
finish()
}
}
targetDeviceMigrationLauncher =
registerForActivityResult(StartMigrationFromTargetDevice(), this::onSignedInResult)
binding.apply {
toolbar.setNavigationOnClickListener {
@@ -94,6 +85,33 @@ class AuthHelpActivity : AuthActivity<ActivityAuthHelpBinding>(ActivityAuthHelpB
}
}
private fun onSignedInResult(result: TargetDeviceMigrationResult?) {
when (result) {
is TargetDeviceMigrationResult.NavigateToSignIn -> {
setResult(RESULT_CANCELED)
finish()
}
is TargetDeviceMigrationResult.SignedIn -> {
setOkResult(AuthHelpResult.SignedInWithEdm(result.userId))
finish()
}
is TargetDeviceMigrationResult.PasswordChangeNeeded -> {
setOkResult(AuthHelpResult.PasswordChangeNeededAfterEdm)
finish()
}
null -> Unit
}
}
private fun setOkResult(result: AuthHelpResult) {
setResult(RESULT_OK, Intent().apply {
putExtra(ARG_RESULT, result)
})
}
companion object {
const val ARG_INPUT = "input"
const val ARG_RESULT = "result"
@@ -31,13 +31,17 @@ import dagger.hilt.android.AndroidEntryPoint
import kotlinx.coroutines.flow.distinctUntilChanged
import kotlinx.coroutines.flow.launchIn
import kotlinx.coroutines.flow.onEach
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.auth.presentation.R
import me.proton.core.auth.presentation.databinding.ActivityChooseAddressBinding
import me.proton.core.auth.presentation.entity.ChooseAddressAuthSecret
import me.proton.core.auth.presentation.entity.ChooseAddressInput
import me.proton.core.auth.presentation.entity.ChooseAddressResult
import me.proton.core.auth.presentation.viewmodel.ChooseAddressViewModel
import me.proton.core.auth.presentation.viewmodel.ChooseAddressViewModel.State
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.domain.entity.UserId
import me.proton.core.observability.domain.metrics.LoginScreenViewTotal
import me.proton.core.network.presentation.util.getUserMessage
@@ -70,6 +74,13 @@ class ChooseAddressActivity : AuthActivity<ActivityChooseAddressBinding>(
requireNotNull(intent?.extras?.getParcelable(ARG_INPUT))
}
private val authSecret: EncryptedAuthSecret by lazy {
when (val s = input.authSecret) {
is ChooseAddressAuthSecret.Passphrase -> EncryptedAuthSecret.Passphrase(s.passphrase)
is ChooseAddressAuthSecret.Password -> EncryptedAuthSecret.Password(s.password)
}
}
private val onBackPressedCallback = object : OnBackPressedCallback(true) {
override fun handleOnBackPressed() {
stopWorkflow()
@@ -142,7 +153,7 @@ class ChooseAddressActivity : AuthActivity<ActivityChooseAddressBinding>(
viewModel.setUsername(
UserId(input.userId),
username = it,
password = input.password,
authSecret = authSecret,
domain = domainInput.text.toString().replace("@", ""),
isTwoPassModeNeeded = input.isTwoPassModeNeeded
)
@@ -181,8 +181,10 @@ class LoginActivity : AuthActivity<ActivityLoginBinding>(ActivityLoginBinding::i
override fun onAuthHelpResult(authHelpResult: AuthHelpResult?) {
super.onAuthHelpResult(authHelpResult)
if (authHelpResult is AuthHelpResult.SignedInWithEdm) {
onSuccess(UserId(authHelpResult.userId))
when (authHelpResult) {
is AuthHelpResult.SignedInWithEdm -> onSuccess(UserId(authHelpResult.userId))
is AuthHelpResult.PasswordChangeNeededAfterEdm -> onChangePassword()
null -> Unit
}
}
@@ -156,8 +156,10 @@ class LoginTwoStepActivity : WebPageListenerActivity(), ProductMetricsDelegateOw
super.onCreate(savedInstanceState)
authHelpLauncher = registerForActivityResult(StartAuthHelp()) { result ->
if (result is AuthHelpResult.SignedInWithEdm) {
onSuccess(UserId(result.userId))
when (result) {
is AuthHelpResult.SignedInWithEdm -> onSuccess(UserId(result.userId))
is AuthHelpResult.PasswordChangeNeededAfterEdm -> onChangePassword()
null -> Unit
}
}
@@ -29,6 +29,7 @@ import kotlinx.coroutines.launch
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.auth.domain.LogTag
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.usecase.AccountAvailability
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup.Result
@@ -110,7 +111,7 @@ class ChooseAddressViewModel @Inject constructor(
fun setUsername(
userId: UserId,
username: String,
password: EncryptedString,
authSecret: EncryptedAuthSecret,
domain: String,
isTwoPassModeNeeded: Boolean
) = viewModelScope.launchWithResultContext {
@@ -120,7 +121,7 @@ class ChooseAddressViewModel @Inject constructor(
flow {
emit(State.Processing)
setupUsername(userId, username)
emit(postLoginSetup(userId, password, domain, isTwoPassModeNeeded))
emit(postLoginSetup(userId, authSecret, domain, isTwoPassModeNeeded))
}.retryOnceWhen(Throwable::primaryKeyExists) {
CoreLogger.e(LogTag.FLOW_ERROR_RETRY, it, "Retrying to upgrade an account")
}.catch { error ->
@@ -159,13 +160,13 @@ class ChooseAddressViewModel @Inject constructor(
private suspend fun postLoginSetup(
userId: UserId,
password: EncryptedString,
authSecret: EncryptedAuthSecret,
domain: String,
isTwoPassModeNeeded: Boolean
): State.AccountSetupResult {
val result = postLoginAccountSetup(
userId = userId,
encryptedPassword = password,
encryptedAuthSecret = authSecret,
requiredAccountType = AccountType.Internal,
isSecondFactorNeeded = false,
isTwoPassModeNeeded = isTwoPassModeNeeded,
@@ -34,6 +34,7 @@ import me.proton.core.account.domain.entity.SessionDetails
import me.proton.core.auth.domain.feature.IsLoginTwoStepEnabled
import me.proton.core.auth.presentation.alert.confirmpass.StartConfirmPassword
import me.proton.core.auth.presentation.entity.AddAccountInput
import me.proton.core.auth.presentation.entity.ChooseAddressAuthSecret
import me.proton.core.auth.presentation.entity.ChooseAddressInput
import me.proton.core.auth.presentation.entity.LoginInput
import me.proton.core.auth.presentation.entity.LoginSsoInput
@@ -285,11 +286,11 @@ class AuthOrchestratorTest {
orchestrator.startChooseAddressWorkflow(account)
}.message
// Then
assertEquals("Password is null for startChooseAddressWorkflow.", message)
assertEquals("Either passphrase or password must be set.", message)
}
@Test
fun `startChooseAddressWorkflow password null`() = runTest {
fun `startChooseAddressWorkflow passphrase and password null`() = runTest {
// Given
orchestrator.register(caller)
val account = mockk<Account>(relaxed = true)
@@ -298,13 +299,14 @@ class AuthOrchestratorTest {
every { account.email } returns "test-email"
every { account.details } returns accountDetails
every { accountDetails.session } returns session
every { session.passphrase } returns null
every { session.password } returns null
// When
val message = assertFailsWith<IllegalStateException> {
orchestrator.startChooseAddressWorkflow(account)
}.message
// Then
assertEquals("Password is null for startChooseAddressWorkflow.", message)
assertEquals("Either passphrase or password must be set.", message)
}
@Test
@@ -336,10 +338,16 @@ class AuthOrchestratorTest {
every { account.userId } returns userId
every { accountDetails.session } returns session
every { session.requiredAccountType } returns AccountType.Internal
every { session.passphrase } returns null
every { session.password } returns encryptedPassword
every { session.twoPassModeEnabled } returns true
// When
val input = ChooseAddressInput(userId.id, encryptedPassword, email, true)
val input = ChooseAddressInput(
userId = userId.id,
authSecret = ChooseAddressAuthSecret.Password(encryptedPassword),
recoveryEmail = email,
isTwoPassModeNeeded = true
)
orchestrator.startChooseAddressWorkflow(account)
// Then
verify(exactly = 1) { chooseAddressLauncher.launch(input) }
@@ -26,6 +26,7 @@ import io.mockk.verify
import kotlinx.coroutines.ExperimentalCoroutinesApi
import kotlinx.coroutines.yield
import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.usecase.AccountAvailability
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.domain.entity.UserId
@@ -245,7 +246,7 @@ class ChooseAddressViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Cor
coEvery {
postLoginAccountSetup.invoke(
userId = any(),
encryptedPassword = any(),
encryptedAuthSecret = any(),
requiredAccountType = any(),
isSecondFactorNeeded = any(),
isTwoPassModeNeeded = any(),
@@ -274,7 +275,7 @@ class ChooseAddressViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Cor
viewModel.setUsername(
userId = userId,
username = "new-username",
password = "password",
authSecret = EncryptedAuthSecret.Password("password"),
domain = "new-domain",
isTwoPassModeNeeded = false
)
@@ -343,7 +344,7 @@ class ChooseAddressViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Cor
viewModel.setUsername(
userId = userId,
username = "username",
password = "password",
authSecret = EncryptedAuthSecret.Password("password"),
domain = "domain",
isTwoPassModeNeeded = false
).join()
@@ -34,6 +34,7 @@ import me.proton.core.auth.domain.usecase.PerformSecondFactor
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.auth.fido.domain.entity.SecondFactorProof
import me.proton.core.auth.presentation.entity.SessionResult
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.domain.entity.UserId
import me.proton.core.network.domain.session.SessionId
import me.proton.core.network.domain.session.SessionProvider
@@ -103,7 +104,7 @@ class SecondFactorViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Unco
coEvery { performSecondFactor.invoke(testSessionId, testSecondFactorProof) } coAnswers {
result("performSecondFactor") { testScopeInfo }
}
coEvery { postLoginAccountSetup.invoke(any(), any(), any(), any(), any(), any()) } returns success
coEvery { postLoginAccountSetup.invoke(any(), any<EncryptedString>(), any(), any(), any(), any()) } returns success
flowTest(viewModel.state) {
// WHEN
viewModel.startSecondFactorFlow(
@@ -133,7 +134,7 @@ class SecondFactorViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Unco
val requiredAccountType = AccountType.Internal
every { testSessionResult.isTwoPassModeNeeded } returns true
coEvery { performSecondFactor.invoke(testSessionId, testSecondFactorProof) } returns testScopeInfo
coEvery { postLoginAccountSetup.invoke(any(), any(), any(), any(), any(), any()) } returns twoPassNeeded
coEvery { postLoginAccountSetup.invoke(any(), any<EncryptedString>(), any(), any(), any(), any()) } returns twoPassNeeded
// WHEN
flowTest(viewModel.state) {
viewModel.startSecondFactorFlow(
@@ -27,6 +27,7 @@ import io.mockk.slot
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.accountmanager.domain.AccountWorkflowHandler
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.crypto.common.keystore.EncryptedString
import me.proton.core.crypto.common.keystore.KeyStoreCrypto
import me.proton.core.domain.entity.UserId
import me.proton.core.test.android.ArchTest
@@ -64,7 +65,7 @@ class TwoPassModeViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Uncon
@Test
fun `mailbox login happy path`() = coroutinesTest {
// GIVEN
coEvery { postLoginAccountSetup.invoke(any(), any(), any(), any(), any(), any()) } returns success
coEvery { postLoginAccountSetup.invoke(any(), any<EncryptedString>(), any(), any(), any(), any()) } returns success
viewModel.state.test {
// WHEN
viewModel.tryUnlockUser(testUserId, testPassword, accountType)
@@ -84,7 +85,7 @@ class TwoPassModeViewModelTest : ArchTest by ArchTest(), CoroutinesTest by Uncon
coEvery {
postLoginAccountSetup.invoke(
any(),
any(),
any<EncryptedString>(),
any(),
any(),
any(),
File diff suppressed because it is too large Load Diff
@@ -202,7 +202,7 @@ abstract class AppDatabase :
companion object {
const val name = "db-account-manager"
const val version = 60
const val version = 61
val migrations = listOf(
AppDatabaseMigrations.MIGRATION_1_2,
@@ -264,6 +264,7 @@ abstract class AppDatabase :
AppDatabaseMigrations.MIGRATION_57_58,
AppDatabaseMigrations.MIGRATION_58_59,
AppDatabaseMigrations.MIGRATION_59_60,
AppDatabaseMigrations.MIGRATION_60_61,
)
fun buildDatabase(context: Context): AppDatabase =
@@ -418,4 +418,10 @@ object AppDatabaseMigrations {
UserSettingsDatabase.MIGRATION_8.migrate(db)
}
}
val MIGRATION_60_61 = object : Migration(60, 61) {
override fun migrate(db: SupportSQLiteDatabase) {
AccountDatabase.MIGRATION_10.migrate(db)
}
}
}
+1
View File
@@ -21,6 +21,7 @@ import studio.forface.easygradle.dsl.android.*
plugins {
protonAndroidLibrary
id("kotlin-parcelize")
}
protonBuild {
@@ -18,8 +18,10 @@
package me.proton.core.crypto.android.keystore
import android.os.Parcel
import androidx.room.RoomDatabase
import androidx.room.TypeConverter
import kotlinx.parcelize.Parceler
import me.proton.core.crypto.common.keystore.EncryptedByteArray
/**
@@ -35,3 +37,17 @@ class CryptoConverters {
fun fromByteArrayToEncryptedByteArray(value: ByteArray?): EncryptedByteArray? =
value?.let { EncryptedByteArray(it) }
}
object EncryptedByteArrayParceler : Parceler<EncryptedByteArray> {
override fun EncryptedByteArray.write(parcel: Parcel, flags: Int) {
parcel.writeInt(array.size)
parcel.writeByteArray(array)
}
override fun create(parcel: Parcel): EncryptedByteArray {
val size = parcel.readInt()
val buf = ByteArray(size)
parcel.readByteArray(buf)
return EncryptedByteArray(buf)
}
}
@@ -26,13 +26,20 @@ import dagger.hilt.android.components.ViewModelComponent
import dagger.hilt.android.scopes.ViewModelScoped
import dagger.hilt.components.SingletonComponent
import me.proton.core.devicemigration.data.feature.IsEasyDeviceMigrationEnabledImpl
import me.proton.core.devicemigration.data.usecase.IsEasyDeviceMigrationAvailableImpl
import me.proton.core.devicemigration.domain.feature.IsEasyDeviceMigrationEnabled
import me.proton.core.devicemigration.domain.usecase.GenerateEdmCode
import me.proton.core.devicemigration.domain.usecase.IsEasyDeviceMigrationAvailable
import me.proton.core.devicemigration.domain.usecase.ObserveEdmCode
@Module
@InstallIn(SingletonComponent::class)
public interface CoreDeviceMigrationModule
public interface CoreDeviceMigrationModule {
@Binds
public fun bindIsEasyDeviceMigrationAvailable(
impl: IsEasyDeviceMigrationAvailableImpl
): IsEasyDeviceMigrationAvailable
}
@Module
@InstallIn(ViewModelComponent::class)
+6 -1
View File
@@ -35,11 +35,16 @@ android {
dependencies {
api(
project(Module.biometricData),
project(Module.biometricDomain),
project(Module.deviceMigrationDomain),
project(Module.featureFlagDomain)
project(Module.domain),
project(Module.featureFlagDomain),
project(Module.userSettingsDomain),
)
testImplementation(
`coroutines-test`,
`kotlin-test`,
mockk,
)
@@ -0,0 +1,73 @@
/*
* Copyright (c) 2025 Proton AG
* This file is part of Proton AG and ProtonCore.
*
* ProtonCore is free software: you can redistribute it and/or modify
* it under the terms of the GNU General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* ProtonCore is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU General Public License for more details.
*
* You should have received a copy of the GNU General Public License
* along with ProtonCore. If not, see <https://www.gnu.org/licenses/>.
*/
package me.proton.core.devicemigration.data.usecase
import me.proton.core.biometric.data.StrongAuthenticatorsResolver
import me.proton.core.biometric.domain.CheckBiometricAuthAvailability
import me.proton.core.devicemigration.domain.feature.IsEasyDeviceMigrationEnabled
import me.proton.core.devicemigration.domain.usecase.IsEasyDeviceMigrationAvailable
import me.proton.core.domain.entity.Product
import me.proton.core.domain.entity.UserId
import me.proton.core.user.domain.UserManager
import me.proton.core.user.domain.extension.hasKeys
import me.proton.core.user.domain.repository.PassphraseRepository
import me.proton.core.usersettings.domain.usecase.IsUserSettingEnabled
import javax.inject.Inject
public class IsEasyDeviceMigrationAvailableImpl @Inject constructor(
private val checkBiometricAuthAvailability: CheckBiometricAuthAvailability,
private val isEasyDeviceMigrationEnabled: IsEasyDeviceMigrationEnabled,
private val isUserSettingsEnabled: IsUserSettingEnabled,
private val passphraseRepository: PassphraseRepository,
private val product: Product,
private val strongAuthenticatorsResolver: StrongAuthenticatorsResolver,
private val userManager: UserManager
) : IsEasyDeviceMigrationAvailable {
/**
* @param userId If it's not null, the check is performed for the given user as an origin device.
* Otherwise if it's null, the check is performed as a target device (the one trying to log in).
*/
override suspend operator fun invoke(userId: UserId?): Boolean =
isEasyDeviceMigrationEnabled(userId) && when {
userId != null -> isAllowedForUser(userId)
else -> true
}
private suspend fun isAllowedForUser(userId: UserId): Boolean = when {
userId.hasOptedOut() -> false
!hasBiometrics() -> false
// If there's no passphrase, we only allow for VPN users without keys
// (the user may want to migrate from VPN to VPN, so the passphrase is not required then).
!userId.hasPassphrase() -> product == Product.Vpn && !userId.hasKeys()
else -> true
}
private fun hasBiometrics(): Boolean =
checkBiometricAuthAvailability(authenticatorsResolver = strongAuthenticatorsResolver).canAttemptBiometricAuth()
private suspend fun UserId.hasOptedOut(): Boolean = runCatching {
isUserSettingsEnabled(this) { easyDeviceMigrationOptOut }
}.getOrNull() ?: true
private suspend fun UserId.hasPassphrase(): Boolean =
passphraseRepository.getPassphrase(this) != null
private suspend fun UserId.hasKeys(): Boolean =
userManager.getUser(this).hasKeys()
}
@@ -16,7 +16,7 @@
* along with ProtonCore. If not, see <https://www.gnu.org/licenses/>.
*/
package me.proton.core.devicemigration.domain.usecase
package me.proton.core.devicemigration.data.usecase
import io.mockk.MockKAnnotations
import io.mockk.coEvery
@@ -24,8 +24,13 @@ import io.mockk.every
import io.mockk.impl.annotations.MockK
import io.mockk.mockk
import kotlinx.coroutines.test.runTest
import me.proton.core.biometric.data.StrongAuthenticatorsResolver
import me.proton.core.biometric.domain.CheckBiometricAuthAvailability
import me.proton.core.devicemigration.domain.feature.IsEasyDeviceMigrationEnabled
import me.proton.core.devicemigration.domain.usecase.IsEasyDeviceMigrationAvailable
import me.proton.core.domain.entity.Product
import me.proton.core.domain.entity.UserId
import me.proton.core.user.domain.UserManager
import me.proton.core.user.domain.repository.PassphraseRepository
import me.proton.core.usersettings.domain.usecase.IsUserSettingEnabled
import kotlin.test.BeforeTest
@@ -34,6 +39,9 @@ import kotlin.test.assertFalse
import kotlin.test.assertTrue
class IsEasyDeviceMigrationAvailableImplTest {
@MockK
private lateinit var checkBiometricAuthAvailability: CheckBiometricAuthAvailability
@MockK
private lateinit var isEasyDeviceMigrationEnabled: IsEasyDeviceMigrationEnabled
@@ -43,21 +51,23 @@ class IsEasyDeviceMigrationAvailableImplTest {
@MockK
private lateinit var passphraseRepository: PassphraseRepository
@MockK
private lateinit var strongAuthenticatorsResolver: StrongAuthenticatorsResolver
@MockK
private lateinit var userManager: UserManager
private lateinit var tested: IsEasyDeviceMigrationAvailable
@BeforeTest
fun setUp() {
MockKAnnotations.init(this)
tested = IsEasyDeviceMigrationAvailable(
isEasyDeviceMigrationEnabled = isEasyDeviceMigrationEnabled,
isUserSettingsEnabled = isUserSettingEnabled,
passphraseRepository = passphraseRepository
)
}
@Test
fun `edm feature flag disabled`() = runTest {
// Given
makeTested()
every { isEasyDeviceMigrationEnabled(any()) } returns false
// When
@@ -70,6 +80,7 @@ class IsEasyDeviceMigrationAvailableImplTest {
@Test
fun `edm user setting disabled`() = runTest {
// Given
makeTested()
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns true // easyDeviceMigrationOptOut
@@ -83,6 +94,8 @@ class IsEasyDeviceMigrationAvailableImplTest {
@Test
fun `edm user setting enabled`() = runTest {
// Given
makeTested()
every { checkBiometricAuthAvailability(any(), any()) } returns CheckBiometricAuthAvailability.Result.Success
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns false // easyDeviceMigrationOptOut
coEvery { passphraseRepository.getPassphrase(any()) } returns mockk()
@@ -97,6 +110,8 @@ class IsEasyDeviceMigrationAvailableImplTest {
@Test
fun `edm user setting disabled because no passphrase`() = runTest {
// Given
makeTested()
every { checkBiometricAuthAvailability(any(), any()) } returns CheckBiometricAuthAvailability.Result.Success
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns false // easyDeviceMigrationOptOut
coEvery { passphraseRepository.getPassphrase(any()) } returns null
@@ -108,9 +123,48 @@ class IsEasyDeviceMigrationAvailableImplTest {
assertFalse(result)
}
@Test
fun `edm user setting enabled on vpn with no passphrase and no user keys`() = runTest {
// Given
makeTested(Product.Vpn)
every { checkBiometricAuthAvailability(any(), any()) } returns CheckBiometricAuthAvailability.Result.Success
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns false // easyDeviceMigrationOptOut
coEvery { passphraseRepository.getPassphrase(any()) } returns null
coEvery { userManager.getUser(any()) } returns mockk(relaxed = true) {
every { keys } returns emptyList()
}
// When
val result = tested(userId = UserId("id-1"))
// Then
assertTrue(result)
}
@Test
fun `edm user setting disabled on vpn with no passphrase and user keys`() = runTest {
// Given
makeTested(Product.Vpn)
every { checkBiometricAuthAvailability(any(), any()) } returns CheckBiometricAuthAvailability.Result.Success
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns false // easyDeviceMigrationOptOut
coEvery { passphraseRepository.getPassphrase(any()) } returns null
coEvery { userManager.getUser(any()) } returns mockk(relaxed = true) {
every { keys } returns listOf(mockk())
}
// When
val result = tested(userId = UserId("id-1"))
// Then
assertFalse(result)
}
@Test
fun `edm enabled for anonymous user`() = runTest {
// Given
makeTested()
every { isEasyDeviceMigrationEnabled(any()) } returns true
// When
@@ -119,4 +173,33 @@ class IsEasyDeviceMigrationAvailableImplTest {
// Then
assertTrue(result)
}
@Test
fun `edm user setting disabled because no biometrics`() = runTest {
// Given
makeTested()
every {
checkBiometricAuthAvailability(any(), any())
} returns CheckBiometricAuthAvailability.Result.Failure.NotEnrolled
every { isEasyDeviceMigrationEnabled(any()) } returns true
coEvery { isUserSettingEnabled(any(), any(), any()) } returns false // easyDeviceMigrationOptOut
// When
val result = tested(userId = UserId("id-1"))
// Then
assertFalse(result)
}
private fun makeTested(product: Product = Product.Mail) {
tested = IsEasyDeviceMigrationAvailableImpl(
checkBiometricAuthAvailability = checkBiometricAuthAvailability,
isEasyDeviceMigrationEnabled = isEasyDeviceMigrationEnabled,
isUserSettingsEnabled = isUserSettingEnabled,
passphraseRepository = passphraseRepository,
product = product,
strongAuthenticatorsResolver = strongAuthenticatorsResolver,
userManager = userManager
)
}
}
+2 -1
View File
@@ -26,8 +26,9 @@ publishOption.shouldBePublishedAsLib = true
dependencies {
api(
project(Module.authDomain),
project(Module.domain),
project(Module.userSettingsDomain),
project(Module.networkDomain),
)
implementation(
@@ -18,27 +18,12 @@
package me.proton.core.devicemigration.domain.usecase
import me.proton.core.devicemigration.domain.feature.IsEasyDeviceMigrationEnabled
import me.proton.core.domain.entity.UserId
import me.proton.core.user.domain.repository.PassphraseRepository
import me.proton.core.usersettings.domain.usecase.IsUserSettingEnabled
import javax.inject.Inject
public class IsEasyDeviceMigrationAvailable @Inject constructor(
private val isEasyDeviceMigrationEnabled: IsEasyDeviceMigrationEnabled,
private val isUserSettingsEnabled: IsUserSettingEnabled,
private val passphraseRepository: PassphraseRepository
) {
public suspend operator fun invoke(userId: UserId?): Boolean =
isEasyDeviceMigrationEnabled(userId) && when {
userId != null -> !userId.hasOptedOut() && userId.hasPassphrase()
else -> true
}
private suspend fun UserId.hasOptedOut(): Boolean = runCatching {
isUserSettingsEnabled(this) { easyDeviceMigrationOptOut }
}.getOrNull() ?: true
private suspend fun UserId.hasPassphrase(): Boolean =
passphraseRepository.getPassphrase(this) != null
public interface IsEasyDeviceMigrationAvailable {
/**
* @param userId If it's not null, the check is performed for the given user as an origin device.
* Otherwise if it's null, the check is performed as a target device (the one trying to log in).
*/
public suspend operator fun invoke(userId: UserId?): Boolean
}
@@ -80,7 +80,7 @@ public class PullEdmSessionFork @Inject constructor(
public data object Awaiting : Result
public data object Loading : Result
public data class Success(
val passphrase: EncryptedByteArray,
val passphrase: EncryptedByteArray?,
val session: Session.Authenticated
) : Result
@@ -31,6 +31,7 @@ android {
namespace = "me.proton.core.devicemigration.presentation"
buildFeatures {
buildConfig = true
resValues = true
viewBinding = true
}
@@ -65,9 +65,12 @@ public class StartMigrationFromTargetDevice : ActivityResultContract<Unit, Targe
@Parcelize
public sealed interface TargetDeviceMigrationResult : Parcelable {
@Parcelize
/** It was not possible to sign in, and the user should be navigated to the sign-in screen. */
public data object NavigateToSignIn : TargetDeviceMigrationResult
@Parcelize
/** User was signed in successfully. */
public data class SignedIn(val userId: String) : TargetDeviceMigrationResult
/** User was signed in, but the password needs to be changed. */
public data object PasswordChangeNeeded : TargetDeviceMigrationResult
}
@@ -53,15 +53,26 @@ public class TargetDeviceMigrationActivity : ProtonActivity() {
},
onNavigateBack = { finish() },
onSuccess = { userId: UserId ->
setResult(RESULT_OK, Intent().apply {
putExtra(ARG_RESULT, TargetDeviceMigrationResult.SignedIn(userId.id))
})
finish()
}
onSuccess(userId, shouldChangePassword = false)
},
onSuccessAndPasswordChange = { userId: UserId ->
onSuccess(userId, shouldChangePassword = true)
},
)
}
}
private fun onSuccess(userId: UserId, shouldChangePassword: Boolean) {
val result = when {
shouldChangePassword -> TargetDeviceMigrationResult.PasswordChangeNeeded
else -> TargetDeviceMigrationResult.SignedIn(userId.id)
}
setResult(RESULT_OK, Intent().apply {
putExtra(ARG_RESULT, result)
})
finish()
}
internal companion object {
const val ARG_RESULT = "result"
}
@@ -35,6 +35,7 @@ internal object TargetDeviceMigrationRoutes {
onBackToSignIn: () -> Unit = {},
onNavigateBack: () -> Unit = {},
onSuccess: (userId: UserId) -> Unit,
onSuccessAndPasswordChange: (userId: UserId) -> Unit,
) {
composable(
route = Route.SignIn.Deeplink
@@ -42,7 +43,8 @@ internal object TargetDeviceMigrationRoutes {
SignInScreen(
onBackToSignIn = onBackToSignIn,
onNavigateBack = onNavigateBack,
onSuccess = onSuccess
onSuccess = onSuccess,
onSuccessAndPasswordChange = onSuccessAndPasswordChange,
)
}
}
@@ -29,7 +29,6 @@ import kotlinx.coroutines.flow.flow
import me.proton.core.biometric.data.StrongAuthenticatorsResolver
import me.proton.core.biometric.domain.BiometricAuthErrorCode
import me.proton.core.biometric.domain.BiometricAuthResult
import me.proton.core.biometric.domain.CheckBiometricAuthAvailability
import me.proton.core.compose.effect.Effect
import me.proton.core.compose.viewmodel.BaseViewModel
import me.proton.core.devicemigration.domain.usecase.DecodeEdmCode
@@ -43,7 +42,6 @@ import javax.inject.Inject
@HiltViewModel
internal class SignInIntroViewModel @Inject constructor(
private val checkBiometricAuthAvailability: CheckBiometricAuthAvailability,
@ApplicationContext private val context: Context,
private val decodeEdmCode: DecodeEdmCode,
private val pushEdmSessionFork: PushEdmSessionFork,
@@ -71,13 +69,7 @@ internal class SignInIntroViewModel @Inject constructor(
}
private fun onStart() = flow {
when {
shouldStartBiometricsCheck() ->
emit(idleWithEffect(SignInIntroEvent.LaunchBiometricsCheck(strongAuthenticatorsResolver)))
else ->
emit(idleWithEffect(SignInIntroEvent.LaunchQrScanner))
}
emit(idleWithEffect(SignInIntroEvent.LaunchBiometricsCheck(strongAuthenticatorsResolver)))
}
private fun onBiometricAuthResult(result: BiometricAuthResult) = flow {
@@ -122,9 +114,6 @@ internal class SignInIntroViewModel @Inject constructor(
private fun stateWithEffect(state: SignInIntroState, event: SignInIntroEvent): SignInIntroStateHolder =
SignInIntroStateHolder(Effect.of(event), state)
private fun shouldStartBiometricsCheck() =
checkBiometricAuthAvailability(authenticatorsResolver = strongAuthenticatorsResolver).canAttemptBiometricAuth()
}
private val BiometricAuthErrorCode.shouldDisplayErrorMessage: Boolean
@@ -22,4 +22,5 @@ import me.proton.core.domain.entity.UserId
internal sealed interface SignInEvent {
data class SignedIn(val userId: UserId) : SignInEvent
data class SignedInAndPasswordChange(val userId: UserId) : SignInEvent
}
@@ -25,5 +25,5 @@ internal sealed interface SignInOperation
internal sealed interface SignInAction : SignInOperation {
data class Load(val unused: Long = System.currentTimeMillis()) : SignInAction
data class SessionForkPulled(val passphrase: EncryptedByteArray, val session: Session.Authenticated) : SignInAction
data class SessionForkPulled(val passphrase: EncryptedByteArray?, val session: Session.Authenticated) : SignInAction
}
@@ -52,6 +52,7 @@ import androidx.compose.ui.text.style.TextAlign
import androidx.compose.ui.tooling.preview.Preview
import androidx.compose.ui.unit.dp
import androidx.compose.ui.unit.times
import androidx.core.graphics.createBitmap
import androidx.lifecycle.compose.collectAsStateWithLifecycle
import me.proton.core.compose.component.ProtonBackButton
import me.proton.core.compose.component.ProtonSolidButton
@@ -74,32 +75,33 @@ internal fun SignInScreen(
onBackToSignIn: () -> Unit,
onNavigateBack: () -> Unit,
onSuccess: (userId: UserId) -> Unit,
onSuccessAndPasswordChange: (userId: UserId) -> Unit,
modifier: Modifier = Modifier,
viewModel: SignInViewModel? = hiltViewModelOrNull<SignInViewModel>()
) {
val state by viewModel?.state?.collectAsStateWithLifecycle()
?: remember { derivedStateOf { SignInStateHolder(state = SignInState.Loading) } }
?: remember { derivedStateOf { SignInState.Loading } }
SignInScreen(
state = state.state,
effect = state.effect,
state = state,
modifier = modifier,
onBackToSignIn = onBackToSignIn,
onNavigateBack = onNavigateBack,
onSuccess = onSuccess
onSuccess = onSuccess,
onSuccessAndPasswordChange = onSuccessAndPasswordChange,
)
}
@Composable
internal fun SignInScreen(
state: SignInState,
effect: Effect<SignInEvent>?,
modifier: Modifier = Modifier,
onBackToSignIn: () -> Unit = {},
onNavigateBack: () -> Unit = {},
onSuccess: (userId: UserId) -> Unit = {},
onSuccessAndPasswordChange: (userId: UserId) -> Unit = {},
) {
val title = when (state) {
is SignInState.UnrecoverableError -> ""
is SignInState.Failure -> ""
else -> stringResource(R.string.target_sign_in_title)
}
Scaffold(
@@ -107,7 +109,7 @@ internal fun SignInScreen(
topBar = { SignInTopBar(onBackClicked = onNavigateBack, title = title) }
) { padding ->
Box(modifier = modifier.padding(padding)) {
if (state is SignInState.UnrecoverableError) {
if (state is SignInState.Failure) {
SignInErrorContent(state, onBackToSignIn = onBackToSignIn)
} else {
SignInContent(state)
@@ -115,7 +117,11 @@ internal fun SignInScreen(
}
}
SignInEffects(effect, onSuccess)
SignInEffects(
effect = state.effect,
onSuccess = onSuccess,
onSuccessAndPasswordChange = onSuccessAndPasswordChange,
)
}
@Composable
@@ -173,14 +179,14 @@ private fun SignInContent(
.size(qrBitmapSize)
)
is SignInState.UnrecoverableError -> Image(
is SignInState.Failure -> Image(
painter = painterResource(R.drawable.ic_proton_cross_big),
contentDescription = null,
modifier = Modifier.align(Alignment.Center),
colorFilter = ColorFilter.tint(Color.Black)
)
SignInState.SuccessfullySignedIn -> Image(
is SignInState.SuccessfullySignedIn -> Image(
painter = painterResource(R.drawable.ic_proton_checkmark),
contentDescription = null,
modifier = Modifier.align(Alignment.Center),
@@ -204,7 +210,7 @@ private fun SignInContent(
@Composable
private fun SignInErrorContent(
state: SignInState.UnrecoverableError,
state: SignInState.Failure,
onBackToSignIn: () -> Unit,
modifier: Modifier = Modifier
) {
@@ -228,7 +234,7 @@ private fun SignInErrorContent(
)
Text(
text = stringResource(R.string.target_sign_in_error_description),
text = state.message,
style = LocalTypography.current.body2Regular,
textAlign = TextAlign.Center,
modifier = Modifier
@@ -238,12 +244,14 @@ private fun SignInErrorContent(
Spacer(modifier = Modifier.weight(1.0f))
ProtonSolidButton(
onClick = state.onRetry,
contained = false,
modifier = Modifier.heightIn(min = ProtonDimens.DefaultButtonMinHeight)
) {
Text(text = stringResource(R.string.target_sign_in_error_new_qr))
if (state.onRetry != null) {
ProtonSolidButton(
onClick = state.onRetry,
contained = false,
modifier = Modifier.heightIn(min = ProtonDimens.DefaultButtonMinHeight)
) {
Text(text = stringResource(R.string.target_sign_in_error_new_qr))
}
}
ProtonTextButton(
@@ -262,11 +270,13 @@ private fun SignInErrorContent(
private fun SignInEffects(
effect: Effect<SignInEvent>?,
onSuccess: (userId: UserId) -> Unit,
onSuccessAndPasswordChange: (userId: UserId) -> Unit = {},
) {
LaunchedEffect(effect) {
effect?.consume { event ->
when (event) {
is SignInEvent.SignedIn -> onSuccess(event.userId)
is SignInEvent.SignedInAndPasswordChange -> onSuccessAndPasswordChange(event.userId)
}
}
}
@@ -278,8 +288,15 @@ private fun SignInEffects(
private fun SignInScreenPreview() {
ProtonTheme {
SignInScreen(
state = SignInState.Loading,
effect = null
state = SignInState.Idle(
"qr-code",
generateBitmap = { _, size ->
createBitmap(
size.value.toInt(),
size.value.toInt(),
Bitmap.Config.ARGB_8888
)
})
)
}
}
@@ -290,8 +307,7 @@ private fun SignInScreenPreview() {
private fun SignInScreenErrorPreview() {
ProtonTheme {
SignInScreen(
state = SignInState.UnrecoverableError(onRetry = {}),
effect = null
state = SignInState.Failure(message = "Error", onRetry = {})
)
}
}
@@ -22,14 +22,9 @@ import android.graphics.Bitmap
import androidx.compose.ui.unit.Dp
import me.proton.core.compose.effect.Effect
internal data class SignInStateHolder(
val effect: Effect<SignInEvent>? = null,
val state: SignInState
)
internal sealed interface SignInState {
data object Loading : SignInState
data class Idle(val qrCode: String, val generateBitmap: suspend (String, Dp) -> Bitmap) : SignInState
data class UnrecoverableError(val onRetry: () -> Unit) : SignInState
data object SuccessfullySignedIn : SignInState
internal sealed class SignInState(open val effect: Effect<SignInEvent>? = null) {
data object Loading : SignInState()
data class Idle(val qrCode: String, val generateBitmap: suspend (String, Dp) -> Bitmap) : SignInState()
data class Failure(val message: String, val onRetry: (() -> Unit)?) : SignInState()
data class SuccessfullySignedIn(override val effect: Effect<SignInEvent>) : SignInState(effect)
}
@@ -18,76 +18,122 @@
package me.proton.core.devicemigration.presentation.signin
import android.content.Context
import dagger.hilt.android.lifecycle.HiltViewModel
import kotlinx.coroutines.delay
import dagger.hilt.android.qualifiers.ApplicationContext
import kotlinx.coroutines.flow.Flow
import kotlinx.coroutines.flow.FlowCollector
import kotlinx.coroutines.flow.flatMapLatest
import kotlinx.coroutines.flow.flow
import kotlinx.coroutines.flow.map
import kotlinx.coroutines.flow.onStart
import me.proton.core.account.domain.entity.AccountType
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.usecase.CreateLoginSessionFromFork
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.compose.effect.Effect
import me.proton.core.compose.viewmodel.BaseViewModel
import me.proton.core.devicemigration.domain.usecase.ObserveEdmCode
import me.proton.core.devicemigration.domain.usecase.PullEdmSessionFork
import me.proton.core.devicemigration.presentation.BuildConfig
import me.proton.core.devicemigration.presentation.R
import me.proton.core.devicemigration.presentation.qr.QrBitmapGenerator
import me.proton.core.util.kotlin.CoreLogger
import javax.inject.Inject
import kotlin.time.Duration.Companion.seconds
@HiltViewModel
internal class SignInViewModel @Inject constructor(
private val accountType: AccountType,
@ApplicationContext private val context: Context,
private val createLoginSessionFromFork: CreateLoginSessionFromFork,
private val observeEdmCode: ObserveEdmCode,
private val postLoginAccountSetup: PostLoginAccountSetup,
private val pullEdmSessionFork: PullEdmSessionFork,
private val qrBitmapGenerator: QrBitmapGenerator,
) : BaseViewModel<SignInAction, SignInStateHolder>(
) : BaseViewModel<SignInAction, SignInState>(
initialAction = SignInAction.Load(),
initialState = SignInStateHolder(state = SignInState.Loading),
initialState = SignInState.Loading
) {
override fun onAction(action: SignInAction): Flow<SignInStateHolder> = when (action) {
override fun onAction(action: SignInAction): Flow<SignInState> = when (action) {
is SignInAction.Load -> onLoad()
is SignInAction.SessionForkPulled -> onSessionForkPulled(action)
}
override suspend fun FlowCollector<SignInStateHolder>.onError(throwable: Throwable) {
emit(stateWithUnrecoverableError())
override suspend fun FlowCollector<SignInState>.onError(throwable: Throwable) {
emit(stateWithUnrecoverableError(onRetry = { perform(SignInAction.Load()) }))
}
private fun onLoad(): Flow<SignInStateHolder> = observeEdmCode(sessionId = null).flatMapLatest { edmCodeResult ->
private fun onLoad(): Flow<SignInState> = observeEdmCode(sessionId = null).flatMapLatest { edmCodeResult ->
pullEdmSessionFork(edmCodeResult.edmParams.encryptionKey, edmCodeResult.selector).map { pullResult ->
Pair(pullResult, edmCodeResult.qrCodeContent)
}
}.map { (pullResult, qrCodeContent) ->
if (BuildConfig.DEBUG) {
CoreLogger.d("GenerateEdmCode", "QR code: $qrCodeContent")
}
when (pullResult) {
is PullEdmSessionFork.Result.Awaiting,
is PullEdmSessionFork.Result.Loading -> SignInStateHolder(
state = SignInState.Idle(
qrCode = qrCodeContent,
generateBitmap = qrBitmapGenerator::invoke
)
is PullEdmSessionFork.Result.Loading -> SignInState.Idle(
qrCode = qrCodeContent,
generateBitmap = qrBitmapGenerator::invoke
)
is PullEdmSessionFork.Result.Success -> {
perform(SignInAction.SessionForkPulled(pullResult.passphrase, pullResult.session))
SignInStateHolder(state = SignInState.Loading)
SignInState.Loading
}
is PullEdmSessionFork.Result.UnrecoverableError -> stateWithUnrecoverableError()
is PullEdmSessionFork.Result.UnrecoverableError ->
stateWithUnrecoverableError(onRetry = { perform(SignInAction.Load()) })
}
}.onStart { emit(SignInStateHolder(state = SignInState.Loading)) }
}.onStart { emit(SignInState.Loading) }
private fun onSessionForkPulled(action: SignInAction.SessionForkPulled) = flow {
emit(SignInStateHolder(state = SignInState.Loading))
delay(3.seconds) // TODO remove
TODO("perform post-login actions")
SignInStateHolder(
effect = Effect.of(SignInEvent.SignedIn(action.session.userId)),
state = SignInState.SuccessfullySignedIn
emit(SignInState.Loading)
createLoginSessionFromFork(accountType, action.passphrase, action.session)
val authSecret = action.passphrase?.let { EncryptedAuthSecret.Passphrase(it) } ?: EncryptedAuthSecret.Absent
val result = postLoginAccountSetup(
userId = action.session.userId,
encryptedAuthSecret = authSecret,
requiredAccountType = accountType,
isSecondFactorNeeded = false,
isTwoPassModeNeeded = false,
temporaryPassword = false,
)
val state = when (result) {
// Most `Result.Need.*` are handled separately by `AccountManagerObserver`.
is PostLoginAccountSetup.Result.Need.ChooseUsername,
is PostLoginAccountSetup.Result.Need.DeviceSecret,
is PostLoginAccountSetup.Result.Need.SecondFactor,
is PostLoginAccountSetup.Result.Need.TwoPassMode,
is PostLoginAccountSetup.Result.AccountReady ->
SignInState.SuccessfullySignedIn(Effect.of(SignInEvent.SignedIn(action.session.userId)))
is PostLoginAccountSetup.Result.Need.ChangePassword ->
SignInState.SuccessfullySignedIn(Effect.of(SignInEvent.SignedInAndPasswordChange(action.session.userId)))
is PostLoginAccountSetup.Result.Error.UnlockPrimaryKeyError -> stateWithUnrecoverableError(
message = context.getString(R.string.target_sign_in_passphrase_error),
onRetry = null
)
is PostLoginAccountSetup.Result.Error.UserCheckError -> stateWithUnrecoverableError(
message = result.error.localizedMessage,
onRetry = null
)
}
emit(state)
}
private fun stateWithUnrecoverableError() = SignInStateHolder(
state = SignInState.UnrecoverableError(
onRetry = { perform(SignInAction.Load()) }
)
private fun stateWithUnrecoverableError(
message: String = context.getString(R.string.target_sign_in_retryable_error),
onRetry: (() -> Unit)?
) = SignInState.Failure(
message = message,
onRetry = onRetry
)
}
@@ -44,7 +44,8 @@
<string name="target_sign_in_scan_code">Scan this code with your phone camera to sign in instantly.</string>
<string name="target_sign_in_instructions">1. Open the Proton app on your phone\n2. Tap into <b>Settings</b>, then tap <b>Sign in to another device</b>\n3. Tap <b>Scan QR code</b></string>
<string name="target_sign_in_error_title">Something went wrong</string>
<string name="target_sign_in_error_description">We couldn\'t sign you in.\nPlease scan a new QR code to try again.</string>
<string name="target_sign_in_retryable_error">We couldn\'t sign you in.\nPlease scan a new QR code to try again.</string>
<string name="target_sign_in_passphrase_error">We couldn\'t sign you in.\nPlease sign in with your password.</string>
<string name="target_sign_in_error_new_qr">New QR code</string>
<string name="target_sign_in_error_back_to_signin">Back to sign-in</string>
</resources>
@@ -13,7 +13,6 @@ import me.proton.core.auth.domain.entity.SessionForkUserCode
import me.proton.core.biometric.data.StrongAuthenticatorsResolver
import me.proton.core.biometric.domain.BiometricAuthErrorCode
import me.proton.core.biometric.domain.BiometricAuthResult
import me.proton.core.biometric.domain.CheckBiometricAuthAvailability
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.devicemigration.domain.entity.ChildClientId
import me.proton.core.devicemigration.domain.entity.EdmParams
@@ -27,13 +26,9 @@ import kotlin.test.BeforeTest
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertIs
import kotlin.test.assertNull
import kotlin.test.assertSame
class SignInIntroViewModelTest : CoroutinesTest by CoroutinesTest() {
@MockK
private lateinit var checkBiometricAuthAvailability: CheckBiometricAuthAvailability
@MockK
private lateinit var context: Context
@@ -55,7 +50,6 @@ class SignInIntroViewModelTest : CoroutinesTest by CoroutinesTest() {
fun setUp() {
MockKAnnotations.init(this)
tested = SignInIntroViewModel(
checkBiometricAuthAvailability = checkBiometricAuthAvailability,
context = context,
decodeEdmCode = decodeEdmCode,
pushEdmSessionFork = pushEdmSessionFork,
@@ -64,30 +58,8 @@ class SignInIntroViewModelTest : CoroutinesTest by CoroutinesTest() {
)
}
@Test
fun `starting the flow with no biometrics available`() = coroutinesTest {
// GIVEN
every { checkBiometricAuthAvailability(any(), any()) } returns
CheckBiometricAuthAvailability.Result.Failure.Unsupported
tested.state.test {
assertInitialState()
// WHEN
tested.perform(SignInIntroAction.Start)
// THEN
val state = awaitItem()
assertIs<SignInIntroEvent.LaunchQrScanner>(state.effect?.peek())
assertEquals(SignInIntroState.Idle, state.state)
}
}
@Test
fun `starting the flow with biometrics available`() = coroutinesTest {
// GIVEN
every { checkBiometricAuthAvailability(any(), any()) } returns CheckBiometricAuthAvailability.Result.Success
tested.state.test {
assertInitialState()
@@ -35,12 +35,22 @@ class SignInScreenTest(deviceConfig: DeviceConfig) {
@Test
fun `loading state`() {
paparazzi.snapshot {
ProtonTheme {
SignInScreen(state = SignInState.Loading)
}
}
}
@Test
fun `idle state`() {
paparazzi.snapshot {
ProtonTheme {
SignInScreen(
state = SignInState.Loading,
effect = null
)
state = SignInState.Idle(
qrCode = "qr-code",
generateBitmap = { _, _ -> throw NotImplementedError() }
))
}
}
}
@@ -49,10 +59,7 @@ class SignInScreenTest(deviceConfig: DeviceConfig) {
fun `unrecoverable error`() {
paparazzi.snapshot {
ProtonTheme {
SignInScreen(
state = SignInState.UnrecoverableError(onRetry = {}),
effect = null
)
SignInScreen(state = SignInState.Failure(message = "Error", onRetry = {}))
}
}
}
@@ -19,17 +19,20 @@
package me.proton.core.devicemigration.presentation.signin
import android.content.Context
import android.content.res.Resources
import app.cash.turbine.test
import io.mockk.MockKAnnotations
import io.mockk.coEvery
import io.mockk.coJustRun
import io.mockk.every
import io.mockk.impl.annotations.MockK
import io.mockk.mockk
import kotlinx.coroutines.flow.flow
import kotlinx.coroutines.flow.flowOf
import me.proton.core.auth.domain.entity.EncryptedAuthSecret
import me.proton.core.auth.domain.entity.SessionForkSelector
import me.proton.core.auth.domain.entity.SessionForkUserCode
import me.proton.core.auth.domain.usecase.CreateLoginSessionFromFork
import me.proton.core.auth.domain.usecase.PostLoginAccountSetup
import me.proton.core.crypto.common.keystore.EncryptedByteArray
import me.proton.core.devicemigration.domain.entity.ChildClientId
import me.proton.core.devicemigration.domain.entity.EdmCodeResult
@@ -38,24 +41,27 @@ import me.proton.core.devicemigration.domain.entity.EncryptionKey
import me.proton.core.devicemigration.domain.usecase.ObserveEdmCode
import me.proton.core.devicemigration.domain.usecase.PullEdmSessionFork
import me.proton.core.devicemigration.presentation.qr.QrBitmapGenerator
import me.proton.core.domain.entity.UserId
import me.proton.core.network.domain.session.Session
import me.proton.core.test.kotlin.CoroutinesTest
import kotlin.test.BeforeTest
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertIs
import kotlin.test.assertNull
class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
@MockK
private lateinit var context: Context
@MockK
private lateinit var resources: Resources
private lateinit var createLoginSessionFromFork: CreateLoginSessionFromFork
@MockK
private lateinit var observeEdmCode: ObserveEdmCode
@MockK
private lateinit var postLoginAccountSetup: PostLoginAccountSetup
@MockK
private lateinit var pullEdmSessionFork: PullEdmSessionFork
@@ -67,17 +73,32 @@ class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
@BeforeTest
fun setUp() {
MockKAnnotations.init(this)
every { resources.getString(any()) } returns "error-from-resources"
every { context.resources } returns resources
every { context.getString(any()) } returns "string-resource"
coEvery { qrBitmapGenerator.invoke(any(), any(), any(), any()) } returns mockk()
tested = SignInViewModel(observeEdmCode, pullEdmSessionFork, qrBitmapGenerator)
tested = SignInViewModel(
accountType = mockk(),
context = context,
createLoginSessionFromFork = createLoginSessionFromFork,
observeEdmCode = observeEdmCode,
postLoginAccountSetup = postLoginAccountSetup,
pullEdmSessionFork = pullEdmSessionFork,
qrBitmapGenerator = qrBitmapGenerator
)
}
@Test
fun `happy path`() = coroutinesTest {
// GIVEN
val passphrase = EncryptedByteArray(byteArrayOf(1, 2, 3))
val session = mockk<Session.Authenticated>()
val testUserId = UserId("user-id")
val session = mockk<Session.Authenticated> {
every { userId } returns testUserId
}
coJustRun { createLoginSessionFromFork(any(), any(), any()) }
coEvery {
postLoginAccountSetup(any(), any<EncryptedAuthSecret>(), any(), any(), any(), any())
} returns PostLoginAccountSetup.Result.AccountReady(testUserId)
every { pullEdmSessionFork(any(), any(), any()) } returns flowOf(
PullEdmSessionFork.Result.Awaiting,
PullEdmSessionFork.Result.Success(passphrase, session)
@@ -97,23 +118,18 @@ class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
// WHEN
tested.state.test {
// THEN
assertEquals(SignInStateHolder(state = SignInState.Loading), awaitItem())
assertEquals(SignInState.Loading, awaitItem())
assertEquals(
SignInStateHolder(
state = SignInState.Idle(
qrCode = "qr-code",
generateBitmap = qrBitmapGenerator::invoke
)
), awaitItem()
SignInState.Idle(qrCode = "qr-code", generateBitmap = qrBitmapGenerator::invoke),
awaitItem()
)
// loading while performing post-login actions:
assertEquals(SignInStateHolder(state = SignInState.Loading), awaitItem())
assertEquals(SignInState.Loading, awaitItem())
// TODO verify logged in
// val (effect, state) = awaitItem()
// val signedInEvent = assertIs<SignInEvent.SignedIn>(effect?.peek())
// assertEquals(session.userId, signedInEvent.userId)
// assertEquals(SignInState.SuccessfullySignedIn, state)
val state = awaitItem()
assertIs<SignInState.SuccessfullySignedIn>(state)
val signedInEvent = assertIs<SignInEvent.SignedIn>(state.effect.peek())
assertEquals(session.userId, signedInEvent.userId)
}
}
@@ -127,10 +143,8 @@ class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
// WHEN
tested.state.test {
// THEN
assertEquals(SignInStateHolder(state = SignInState.Loading), awaitItem())
val (effect, state) = awaitItem()
assertNull(effect)
assertIs<SignInState.UnrecoverableError>(state)
assertEquals(SignInState.Loading, awaitItem())
assertIs<SignInState.Failure>(awaitItem())
}
}
@@ -147,36 +161,48 @@ class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
// WHEN
tested.state.test {
// THEN
assertEquals(SignInStateHolder(state = SignInState.Loading), awaitItem())
val (effect, state) = awaitItem()
assertNull(effect)
assertIs<SignInState.UnrecoverableError>(state)
assertEquals(SignInState.Loading, awaitItem())
assertIs<SignInState.Failure>(awaitItem())
}
}
@Test
fun `awaiting the fork`() = coroutinesTest {
// GIVEN
val testUserId = UserId("user-id")
val session = mockk<Session.Authenticated> {
every { userId } returns testUserId
}
coJustRun { createLoginSessionFromFork(any(), any(), any()) }
every { observeEdmCode(any()) } returns flowOf(
EdmCodeResult(mockk(relaxed = true), "qr-code", SessionForkSelector("selector"))
)
coEvery {
postLoginAccountSetup(any(), any<EncryptedAuthSecret>(), any(), any(), any(), any())
} returns PostLoginAccountSetup.Result.AccountReady(testUserId)
every { pullEdmSessionFork(any(), any(), any()) } returns flowOf(
PullEdmSessionFork.Result.Loading,
PullEdmSessionFork.Result.Awaiting,
PullEdmSessionFork.Result.Loading,
PullEdmSessionFork.Result.Success(mockk(), mockk()),
PullEdmSessionFork.Result.Success(mockk(), session),
)
// WHEN
tested.state.test {
// THEN
assertIs<SignInState.Loading>(awaitItem().state)
assertIs<SignInState.Loading>(awaitItem())
// stays Idle when pullEdmSessionFork returns Loading or Awaiting:
assertIs<SignInState.Idle>(awaitItem().state)
assertIs<SignInState.Idle>(awaitItem())
// back to loading when pullEdmSessionFork returns Success:
assertIs<SignInState.Loading>(awaitItem().state)
assertIs<SignInState.Loading>(awaitItem())
val state = awaitItem()
assertIs<SignInState.SuccessfullySignedIn>(state)
val signedInEvent = assertIs<SignInEvent.SignedIn>(state.effect.peek())
assertEquals(session.userId, signedInEvent.userId)
}
}
@@ -204,17 +230,16 @@ class SignInViewModelTest : CoroutinesTest by CoroutinesTest() {
// WHEN
tested.state.test {
// THEN
assertIs<SignInState.Loading>(awaitItem().state)
assertIs<SignInState.Loading>(awaitItem())
val (effect1, state1) = awaitItem()
assertNull(effect1)
assertIs<SignInState.UnrecoverableError>(state1)
val state1 = awaitItem()
assertIs<SignInState.Failure>(state1)
// WHEN
state1.onRetry()
state1.onRetry?.invoke()
// THEN
assertIs<SignInState.Loading>(awaitItem().state)
assertIs<SignInState.Loading>(awaitItem())
}
}
}
}
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:2a17ff5455af5d620a25d194305d41a707d5a465a0ae0c2982ef62f7f4d4545b
size 26808
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:f9c30f93c35f9cd6ac01509d176795516c1b4a9aa41ba18d40a54972713c2d7c
size 30780
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:58fb2bef92ee6a8f999420967b5c128d3b130a6c4dbfe98531c726fd7f1ecb21
size 26515
oid sha256:e779ca17bc669ee752e7b2d27cf8f6ce170c2b655a6ae847abbdb1b235b54da1
size 19736
@@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:82d4a8475fa84448bcd5d58390d32bd7a205ebdc385b78dbfcf4bfa053a687a0
size 28092
oid sha256:7c8f5cc859ea4a14037344a004c2f90ec4d48b27c0582aed7bb208134010c8ae
size 20558