mirror of
https://github.com/ProtonMail/protoncore_android.git
synced 2026-06-14 09:54:49 +00:00
feat(auth, device-migration)!: Perform post-login steps after obtaining forked session.
MIGRATION: AccountDatabase.MIGRATION_10
This commit is contained in:
+3562
File diff suppressed because it is too large
Load Diff
+2
-1
@@ -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> =
|
||||
|
||||
+6
@@ -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)
|
||||
}
|
||||
|
||||
+4
-1
@@ -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
|
||||
)
|
||||
|
||||
+2
-1
@@ -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 {
|
||||
|
||||
+5
-1
@@ -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
|
||||
}
|
||||
+1
-1
@@ -25,7 +25,7 @@ import kotlinx.serialization.Serializable
|
||||
@Serializable
|
||||
data class SessionForkPayloadWithKey(
|
||||
@SerialName("keyPassword")
|
||||
val keyPassword: String,
|
||||
val keyPassword: String?,
|
||||
|
||||
@EncodeDefault
|
||||
@SerialName("type")
|
||||
|
||||
+1
@@ -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()
|
||||
)
|
||||
|
||||
+71
@@ -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
|
||||
)
|
||||
}
|
||||
}
|
||||
+1
@@ -62,6 +62,7 @@ class CreateLoginSsoSession @Inject constructor(
|
||||
requiredAccountType = requiredAccountType,
|
||||
secondFactorEnabled = sessionInfo.isSecondFactorNeeded,
|
||||
twoPassModeEnabled = sessionInfo.isTwoPassModeNeeded,
|
||||
passphrase = null,
|
||||
password = null,
|
||||
fido2AuthenticationOptionsJson = sessionInfo.getFido2AuthOptions()?.serialize()
|
||||
)
|
||||
|
||||
+46
-16
@@ -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
-5
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
+4
-3
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+9
-9
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+2
-1
@@ -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
|
||||
|
||||
+73
-4
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
+17
-6
@@ -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.
|
||||
*
|
||||
|
||||
+3
-1
@@ -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
|
||||
}
|
||||
|
||||
+11
-1
@@ -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()
|
||||
}
|
||||
|
||||
+29
-11
@@ -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"
|
||||
|
||||
+12
-1
@@ -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
|
||||
)
|
||||
|
||||
+4
-2
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+4
-2
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+5
-4
@@ -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,
|
||||
|
||||
+12
-4
@@ -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) }
|
||||
|
||||
+4
-3
@@ -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()
|
||||
|
||||
+3
-2
@@ -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(
|
||||
|
||||
+3
-2
@@ -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 =
|
||||
|
||||
+6
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ import studio.forface.easygradle.dsl.android.*
|
||||
|
||||
plugins {
|
||||
protonAndroidLibrary
|
||||
id("kotlin-parcelize")
|
||||
}
|
||||
|
||||
protonBuild {
|
||||
|
||||
+16
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+8
-1
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
+73
@@ -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()
|
||||
}
|
||||
+89
-6
@@ -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
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -26,8 +26,9 @@ publishOption.shouldBePublishedAsLib = true
|
||||
|
||||
dependencies {
|
||||
api(
|
||||
project(Module.authDomain),
|
||||
project(Module.domain),
|
||||
project(Module.userSettingsDomain),
|
||||
project(Module.networkDomain),
|
||||
)
|
||||
|
||||
implementation(
|
||||
|
||||
+6
-21
@@ -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
|
||||
}
|
||||
|
||||
+1
-1
@@ -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
|
||||
}
|
||||
|
||||
+5
-2
@@ -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
|
||||
}
|
||||
|
||||
+16
-5
@@ -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"
|
||||
}
|
||||
|
||||
+3
-1
@@ -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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
+1
-12
@@ -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
|
||||
|
||||
+1
@@ -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
|
||||
}
|
||||
|
||||
+1
-1
@@ -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
|
||||
}
|
||||
|
||||
+38
-22
@@ -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 = {})
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
+5
-10
@@ -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)
|
||||
}
|
||||
|
||||
+72
-26
@@ -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>
|
||||
|
||||
-28
@@ -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()
|
||||
|
||||
|
||||
+14
-7
@@ -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 = {}))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+64
-39
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:2a17ff5455af5d620a25d194305d41a707d5a465a0ae0c2982ef62f7f4d4545b
|
||||
size 26808
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:f9c30f93c35f9cd6ac01509d176795516c1b4a9aa41ba18d40a54972713c2d7c
|
||||
size 30780
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:58fb2bef92ee6a8f999420967b5c128d3b130a6c4dbfe98531c726fd7f1ecb21
|
||||
size 26515
|
||||
oid sha256:e779ca17bc669ee752e7b2d27cf8f6ce170c2b655a6ae847abbdb1b235b54da1
|
||||
size 19736
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:82d4a8475fa84448bcd5d58390d32bd7a205ebdc385b78dbfcf4bfa053a687a0
|
||||
size 28092
|
||||
oid sha256:7c8f5cc859ea4a14037344a004c2f90ec4d48b27c0582aed7bb208134010c8ae
|
||||
size 20558
|
||||
|
||||
Reference in New Issue
Block a user