diff --git a/app/src/main/kotlin/com/wire/android/ui/authentication/create/code/CreateAccountCodeViewModel.kt b/app/src/main/kotlin/com/wire/android/ui/authentication/create/code/CreateAccountCodeViewModel.kt index ff2d21ada07..d6c9af2bc6d 100644 --- a/app/src/main/kotlin/com/wire/android/ui/authentication/create/code/CreateAccountCodeViewModel.kt +++ b/app/src/main/kotlin/com/wire/android/ui/authentication/create/code/CreateAccountCodeViewModel.kt @@ -27,9 +27,10 @@ import androidx.lifecycle.ViewModel import androidx.lifecycle.viewModelScope import com.ramcosta.composedestinations.generated.app.navArgs import com.wire.android.BuildConfig -import com.wire.android.di.ClientScopeProvider import com.wire.android.di.DefaultWebSocketEnabledByDefault import com.wire.android.di.KaliumCoreLogic +import com.wire.android.session.AppUserSessionPreparationResult +import com.wire.android.session.UserSessionPreparationGate import com.wire.android.ui.authentication.create.common.CreateAccountFlowType import com.wire.android.ui.authentication.create.common.CreateAccountNavArgs import com.wire.android.ui.authentication.login.email.LoginEmailViewModel.Companion.RESEND_TIMER_DELAY @@ -37,12 +38,14 @@ import com.wire.android.ui.common.textfield.textAsFlow import com.wire.android.ui.registration.code.CreateAccountCodeResult import com.wire.android.util.WillNeverOccurError import com.wire.android.util.ui.CountdownTimer +import com.wire.kalium.common.error.CoreFailure import com.wire.kalium.logic.CoreLogic import com.wire.kalium.logic.configuration.server.ServerConfig import com.wire.kalium.logic.data.session.StoreSessionParam import com.wire.kalium.logic.data.user.UserId import com.wire.kalium.logic.feature.auth.AddAuthenticatedUserUseCase import com.wire.kalium.logic.feature.auth.autoVersioningAuth.AutoVersionAuthScopeUseCase +import com.wire.kalium.logic.feature.UserSessionScope import com.wire.kalium.logic.feature.client.RegisterClientParam import com.wire.kalium.logic.feature.client.RegisterClientResult import com.wire.kalium.logic.feature.register.RegisterParam @@ -59,7 +62,6 @@ class CreateAccountCodeViewModel @AssistedInject constructor( @Assisted savedStateHandle: SavedStateHandle, @KaliumCoreLogic private val coreLogic: CoreLogic, private val addAuthenticatedUser: AddAuthenticatedUserUseCase, - private val clientScopeProviderFactory: ClientScopeProvider.Factory, defaultServerConfig: ServerConfig.Links, @DefaultWebSocketEnabledByDefault private val defaultWebSocketEnabledByDefault: Boolean ) : ViewModel() { @@ -68,6 +70,8 @@ class CreateAccountCodeViewModel @AssistedInject constructor( fun create(savedStateHandle: SavedStateHandle): CreateAccountCodeViewModel } + private val userSessionPreparationGate by lazy { UserSessionPreparationGate(coreLogic) } + val createAccountNavArgs: CreateAccountNavArgs = savedStateHandle.navArgs() val serverConfig: ServerConfig.Links = createAccountNavArgs.customServerConfig ?: defaultServerConfig @@ -205,7 +209,18 @@ class CreateAccountCodeViewModel @AssistedInject constructor( is AddAuthenticatedUserUseCase.Result.Success -> it.userId } } - registerClient(storedUserId, registerParam.password).let { + val sessionScope = when (val preparation = userSessionPreparationGate.prepare(storedUserId)) { + is AppUserSessionPreparationResult.Ready -> preparation.sessionScope + is AppUserSessionPreparationResult.Failed -> { + updateCodeErrorState( + CreateAccountCodeResult.Error.DialogError.GenericError( + CoreFailure.Unknown(IllegalStateException("User session preparation failed: ${preparation.reason}")) + ) + ) + return@launch + } + } + registerClient(sessionScope, registerParam.password).let { when (it) { is RegisterClientResult.Failure -> { updateCodeErrorState(it.toCodeError(storedUserId)) @@ -229,8 +244,8 @@ class CreateAccountCodeViewModel @AssistedInject constructor( codeState = codeState.copy(loading = false, result = codeError) } - private suspend fun registerClient(userId: UserId, password: String) = - clientScopeProviderFactory.create(userId).clientScope.getOrRegister( + private suspend fun registerClient(sessionScope: UserSessionScope, password: String) = + sessionScope.client.getOrRegister( RegisterClientParam( password = password, capabilities = null, diff --git a/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModel.kt b/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModel.kt index ba9e3ec7352..c3c4fdb9b80 100644 --- a/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModel.kt +++ b/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModel.kt @@ -32,6 +32,7 @@ import com.wire.kalium.logic.data.user.UserId import com.wire.kalium.logic.feature.auth.AddAuthenticatedUserUseCase import com.wire.kalium.logic.feature.auth.AuthenticationResult import com.wire.kalium.logic.feature.auth.DomainLookupUseCase +import com.wire.kalium.logic.feature.UserSessionScope import com.wire.kalium.logic.feature.client.RegisterClientResult @Suppress("TooManyFunctions") @@ -60,6 +61,13 @@ open class LoginViewModel( capabilities: List? = null, ): RegisterClientResult = loginExtension.registerClient(userId, password, secondFactorVerificationCode, capabilities) + suspend fun registerClient( + sessionScope: UserSessionScope, + password: String?, + secondFactorVerificationCode: String? = null, + capabilities: List? = null, + ): RegisterClientResult = loginExtension.registerClient(sessionScope, password, secondFactorVerificationCode, capabilities) + internal suspend fun isInitialSyncCompleted(userId: UserId): Boolean = loginExtension.isInitialSyncCompleted(userId) } diff --git a/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModelExtension.kt b/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModelExtension.kt index 1e2d2cec6d0..aaf9d4fd770 100644 --- a/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModelExtension.kt +++ b/app/src/main/kotlin/com/wire/android/ui/authentication/login/LoginViewModelExtension.kt @@ -22,6 +22,8 @@ import com.wire.android.datastore.UserDataStoreProvider import com.wire.android.di.ClientScopeProvider import com.wire.kalium.logic.data.client.ClientCapability import com.wire.kalium.logic.data.user.UserId +import com.wire.kalium.logic.feature.UserSessionScope +import com.wire.kalium.logic.feature.client.GetOrRegisterClientUseCase import com.wire.kalium.logic.feature.client.RegisterClientParam import com.wire.kalium.logic.feature.client.RegisterClientResult import kotlinx.coroutines.flow.first @@ -38,16 +40,35 @@ class LoginViewModelExtension( capabilities: List? = null, ): RegisterClientResult { val clientScope = clientScopeProviderFactory.create(userId).clientScope - return clientScope.getOrRegister( - RegisterClientParam( - password = password, - capabilities = capabilities, - secondFactorVerificationCode = secondFactorVerificationCode, - modelPostfix = if (BuildConfig.PRIVATE_BUILD) " [${BuildConfig.FLAVOR}_${BuildConfig.BUILD_TYPE}]" else null - ) - ) + return registerClient(clientScope.getOrRegister, password, secondFactorVerificationCode, capabilities) } + suspend fun registerClient( + sessionScope: UserSessionScope, + password: String?, + secondFactorVerificationCode: String? = null, + capabilities: List? = null, + ): RegisterClientResult = registerClient( + sessionScope.client.getOrRegister, + password, + secondFactorVerificationCode, + capabilities, + ) + + private suspend fun registerClient( + getOrRegister: GetOrRegisterClientUseCase, + password: String?, + secondFactorVerificationCode: String?, + capabilities: List?, + ): RegisterClientResult = getOrRegister( + RegisterClientParam( + password = password, + capabilities = capabilities, + secondFactorVerificationCode = secondFactorVerificationCode, + modelPostfix = if (BuildConfig.PRIVATE_BUILD) " [${BuildConfig.FLAVOR}_${BuildConfig.BUILD_TYPE}]" else null + ) + ) + internal suspend fun isInitialSyncCompleted(userId: UserId): Boolean = userDataStoreProvider.getOrCreate(userId).initialSyncCompleted.first() } diff --git a/app/src/main/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModel.kt b/app/src/main/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModel.kt index 8e0784d5247..ddb55b9df12 100644 --- a/app/src/main/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModel.kt +++ b/app/src/main/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModel.kt @@ -27,11 +27,14 @@ import androidx.compose.runtime.mutableStateOf import androidx.compose.runtime.setValue import androidx.lifecycle.SavedStateHandle import androidx.lifecycle.viewModelScope +import com.wire.android.appLogger import com.wire.android.datastore.UserDataStoreProvider import com.wire.android.datastore.GlobalDataStore import com.wire.android.di.ClientScopeProvider import com.wire.android.di.DefaultWebSocketEnabledByDefault import com.wire.android.di.KaliumCoreLogic +import com.wire.android.session.AppUserSessionPreparationResult +import com.wire.android.session.UserSessionPreparationGate import com.wire.android.ui.authentication.toBackendConfigUrl import com.wire.android.ui.authentication.login.LoginNavArgs import com.wire.android.ui.authentication.login.LoginSavedInputStore @@ -137,6 +140,7 @@ class LoginEmailViewModel( ) private val preFilledUserIdentifier: PreFilledUserIdentifierType = loginNavArgs.userHandle ?: PreFilledUserIdentifierType.None + private val userSessionPreparationGate by lazy { UserSessionPreparationGate(coreLogic) } val userIdentifierTextState: TextFieldState = TextFieldState() val passwordTextState: TextFieldState = TextFieldState() @@ -327,9 +331,24 @@ class LoginEmailViewModel( } } + val preparedSessionScope = when (val preparation = userSessionPreparationGate.prepare(storedUserId)) { + is AppUserSessionPreparationResult.Ready -> preparation.sessionScope + is AppUserSessionPreparationResult.Failed -> { + restorePreviousSession() + updateEmailFlowState( + LoginState.Error.DialogError.GenericError( + com.wire.kalium.common.error.CoreFailure.Unknown( + IllegalStateException("User session preparation failed: ${preparation.reason}") + ) + ) + ) + return@launch + } + } + withContext(dispatchers.io()) { if (coreLogic.getGlobalScope().validateEmailUseCase(userIdentifierTextState.text.toString())) { - coreLogic.getSessionScope(storedUserId).users.persistSelfUserEmail(userIdentifierTextState.text.toString()) + preparedSessionScope.users.persistSelfUserEmail(userIdentifierTextState.text.toString()) } else { null } @@ -343,7 +362,7 @@ class LoginEmailViewModel( withContext(dispatchers.io()) { registerClient( - userId = storedUserId, + sessionScope = preparedSessionScope, password = passwordTextState.text.toString(), ) }.let { @@ -370,14 +389,28 @@ class LoginEmailViewModel( } private suspend fun revertNewSession() { - loginJobData.value?.newSessionUserId?.let { newSessionUserId -> - // logout to cancel all session-related actions, remove all sensitive data and free up resources - coreLogic.getSessionScope(newSessionUserId).logout(reason = LogoutReason.SELF_HARD_LOGOUT, waitUntilCompletes = true) - // delete the session to make it seem like the session was never logged in - coreLogic.getGlobalScope().deleteSession(newSessionUserId) + val jobData = loginJobData.value + jobData?.newSessionUserId?.let { newSessionUserId -> + when (val preparation = userSessionPreparationGate.prepare(newSessionUserId)) { + is AppUserSessionPreparationResult.Ready -> { + // logout to cancel session actions and remove sensitive data before deleting the session + preparation.sessionScope.logout( + reason = LogoutReason.SELF_HARD_LOGOUT, + waitUntilCompletes = true, + ) + coreLogic.getGlobalScope().deleteSession(newSessionUserId) + } + + is AppUserSessionPreparationResult.Failed -> appLogger.w( + "Login rollback skipped database deletion because session preparation failed: ${preparation.reason}" + ) + } } - // set the previous session back - coreLogic.getGlobalScope().session.updateCurrentSession(loginJobData.value?.previousSessionUserId) + restorePreviousSession(jobData) + } + + private suspend fun restorePreviousSession(jobData: LoginJobData? = loginJobData.value) { + coreLogic.getGlobalScope().session.updateCurrentSession(jobData?.previousSessionUserId) } private suspend fun revertLogin() { diff --git a/app/src/main/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtension.kt b/app/src/main/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtension.kt index a597fe3750a..6387ff6f1b7 100644 --- a/app/src/main/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtension.kt +++ b/app/src/main/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtension.kt @@ -18,6 +18,8 @@ package com.wire.android.ui.authentication.login.sso import com.wire.android.appLogger +import com.wire.android.session.AppUserSessionPreparationResult +import com.wire.android.session.UserSessionPreparationGate import com.wire.kalium.common.error.CoreFailure import com.wire.kalium.logic.CoreLogic import com.wire.kalium.logic.configuration.server.ServerConfig @@ -48,6 +50,7 @@ class LoginSSOViewModelExtension( private val coreLogic: CoreLogic, private val defaultWebSocketEnabledByDefault: Boolean, ) { + private val userSessionPreparationGate by lazy { UserSessionPreparationGate(coreLogic) } suspend fun withAuthenticationScope( serverConfig: ServerConfig.Links, onAuthScopeFailure: (AutoVersionAuthScopeUseCase.Result.Failure) -> Unit, @@ -149,15 +152,27 @@ class LoginSSOViewModelExtension( when (authenticatedUserResult) { AddAuthenticatedUserUseCase.Result.Failure.SsoIdentityChanged -> onSsoIdentityChanged(session) is AddAuthenticatedUserUseCase.Result.Failure -> onAddAuthenticatedUserFailure(authenticatedUserResult) - is AddAuthenticatedUserUseCase.Result.Success -> onSuccess(authenticatedUserResult.userId) + is AddAuthenticatedUserUseCase.Result.Success -> + when (val preparation = userSessionPreparationGate.prepare(authenticatedUserResult.userId)) { + is AppUserSessionPreparationResult.Ready -> onSuccess(authenticatedUserResult.userId) + is AppUserSessionPreparationResult.Failed -> onAddAuthenticatedUserFailure( + AddAuthenticatedUserUseCase.Result.Failure.Generic(preparation.toCoreFailure()) + ) + } } } - @Suppress("TooGenericExceptionCaught") + @Suppress("TooGenericExceptionCaught", "NestedBlockDepth") suspend fun replaceRetainedSsoSession(session: StoreSessionParam): ReplaceRetainedSsoSessionResult = try { val userId = session.accountTokens.userId - coreLogic.getSessionScope(userId).logout( + val retainedScope = when (val preparation = userSessionPreparationGate.prepare(userId)) { + is AppUserSessionPreparationResult.Ready -> preparation.sessionScope + is AppUserSessionPreparationResult.Failed -> return ReplaceRetainedSsoSessionResult.Failure( + AddAuthenticatedUserUseCase.Result.Failure.Generic(preparation.toCoreFailure()) + ) + } + retainedScope.logout( reason = LogoutReason.SELF_HARD_LOGOUT, waitUntilCompletes = true ) @@ -167,8 +182,16 @@ class LoginSSOViewModelExtension( ReplaceRetainedSsoSessionResult.Failure( AddAuthenticatedUserUseCase.Result.Failure.Generic(deleteResult.cause) ) - DeleteSessionUseCase.Result.Success -> - addAuthenticatedUser(session, replace = false).toReplaceRetainedSsoSessionResult() + DeleteSessionUseCase.Result.Success -> when (val addResult = addAuthenticatedUser(session, replace = false)) { + is AddAuthenticatedUserUseCase.Result.Failure -> ReplaceRetainedSsoSessionResult.Failure(addResult) + is AddAuthenticatedUserUseCase.Result.Success -> + when (val preparation = userSessionPreparationGate.prepare(addResult.userId)) { + is AppUserSessionPreparationResult.Ready -> ReplaceRetainedSsoSessionResult.Success(addResult.userId) + is AppUserSessionPreparationResult.Failed -> ReplaceRetainedSsoSessionResult.Failure( + AddAuthenticatedUserUseCase.Result.Failure.Generic(preparation.toCoreFailure()) + ) + } + } } } catch (exception: CancellationException) { throw exception @@ -179,6 +202,9 @@ class LoginSSOViewModelExtension( } } +private fun AppUserSessionPreparationResult.Failed.toCoreFailure(): CoreFailure = + CoreFailure.Unknown(IllegalStateException("User session preparation failed: $reason")) + private fun AddAuthenticatedUserUseCase.Result.toReplaceRetainedSsoSessionResult(): ReplaceRetainedSsoSessionResult = when (this) { is AddAuthenticatedUserUseCase.Result.Failure -> ReplaceRetainedSsoSessionResult.Failure(this) diff --git a/app/src/main/kotlin/com/wire/android/ui/registration/code/CreateAccountVerificationCodeViewModel.kt b/app/src/main/kotlin/com/wire/android/ui/registration/code/CreateAccountVerificationCodeViewModel.kt index acd876d7fd5..747010baa9c 100644 --- a/app/src/main/kotlin/com/wire/android/ui/registration/code/CreateAccountVerificationCodeViewModel.kt +++ b/app/src/main/kotlin/com/wire/android/ui/registration/code/CreateAccountVerificationCodeViewModel.kt @@ -28,19 +28,22 @@ import androidx.lifecycle.viewModelScope import com.ramcosta.composedestinations.generated.app.navArgs import com.wire.android.BuildConfig import com.wire.android.analytics.RegistrationAnalyticsManagerUseCase -import com.wire.android.di.ClientScopeProvider import com.wire.android.di.DefaultWebSocketEnabledByDefault import com.wire.android.di.KaliumCoreLogic import com.wire.android.feature.analytics.model.AnalyticsEvent +import com.wire.android.session.AppUserSessionPreparationResult +import com.wire.android.session.UserSessionPreparationGate import com.wire.android.ui.authentication.create.common.CreateAccountDataNavArgs import com.wire.android.ui.common.textfield.textAsFlow import com.wire.android.util.WillNeverOccurError +import com.wire.kalium.common.error.CoreFailure import com.wire.kalium.logic.CoreLogic import com.wire.kalium.logic.configuration.server.ServerConfig import com.wire.kalium.logic.data.session.StoreSessionParam import com.wire.kalium.logic.data.user.UserId import com.wire.kalium.logic.feature.auth.AddAuthenticatedUserUseCase import com.wire.kalium.logic.feature.auth.autoVersioningAuth.AutoVersionAuthScopeUseCase +import com.wire.kalium.logic.feature.UserSessionScope import com.wire.kalium.logic.feature.client.RegisterClientParam import com.wire.kalium.logic.feature.client.RegisterClientResult import com.wire.kalium.logic.feature.register.RegisterParam @@ -57,7 +60,6 @@ class CreateAccountVerificationCodeViewModel @AssistedInject constructor( @KaliumCoreLogic private val coreLogic: CoreLogic, private val addAuthenticatedUser: AddAuthenticatedUserUseCase, private val registrationAnalyticsManager: RegistrationAnalyticsManagerUseCase, - private val clientScopeProviderFactory: ClientScopeProvider.Factory, defaultServerConfig: ServerConfig.Links, @DefaultWebSocketEnabledByDefault private val defaultWebSocketEnabledByDefault: Boolean, ) : ViewModel() { @@ -66,6 +68,8 @@ class CreateAccountVerificationCodeViewModel @AssistedInject constructor( fun create(savedStateHandle: SavedStateHandle): CreateAccountVerificationCodeViewModel } + private val userSessionPreparationGate by lazy { UserSessionPreparationGate(coreLogic) } + val createAccountNavArgs: CreateAccountDataNavArgs = savedStateHandle.navArgs() val serverConfig: ServerConfig.Links = createAccountNavArgs.customServerConfig ?: defaultServerConfig @@ -123,7 +127,7 @@ class CreateAccountVerificationCodeViewModel @AssistedInject constructor( codeTextState.clearText() } - @Suppress("ComplexMethod") + @Suppress("ComplexMethod", "LongMethod") private fun onCodeContinue() { codeState = codeState.copy(loading = true) viewModelScope.launch { @@ -185,15 +189,27 @@ class CreateAccountVerificationCodeViewModel @AssistedInject constructor( is AddAuthenticatedUserUseCase.Result.Success -> it.userId } } - registerClient(storedUserId, registerParam) + val sessionScope = when (val preparation = userSessionPreparationGate.prepare(storedUserId)) { + is AppUserSessionPreparationResult.Ready -> preparation.sessionScope + is AppUserSessionPreparationResult.Failed -> { + updateCodeErrorState( + CreateAccountCodeResult.Error.DialogError.GenericError( + CoreFailure.Unknown(IllegalStateException("User session preparation failed: ${preparation.reason}")) + ) + ) + return@launch + } + } + registerClient(storedUserId, sessionScope, registerParam) } } private suspend fun registerClient( storedUserId: UserId, + sessionScope: UserSessionScope, registerParam: RegisterParam.PersonalAccount ) { - registerClient(storedUserId, registerParam.password).let { + registerClient(sessionScope, registerParam.password).let { when (it) { is RegisterClientResult.Failure -> { updateCodeErrorState(it.toCodeError(storedUserId)) @@ -217,8 +233,8 @@ class CreateAccountVerificationCodeViewModel @AssistedInject constructor( codeState = codeState.copy(loading = false, result = codeError) } - private suspend fun registerClient(userId: UserId, password: String) = - clientScopeProviderFactory.create(userId).clientScope.getOrRegister( + private suspend fun registerClient(sessionScope: UserSessionScope, password: String) = + sessionScope.client.getOrRegister( RegisterClientParam( password = password, capabilities = null, diff --git a/app/src/test/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModelTest.kt b/app/src/test/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModelTest.kt index 2bcc5d0c6f8..630bab3c326 100644 --- a/app/src/test/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModelTest.kt +++ b/app/src/test/kotlin/com/wire/android/ui/authentication/login/email/LoginEmailViewModelTest.kt @@ -16,7 +16,7 @@ * along with this program. If not, see http://www.gnu.org/licenses/. */ -@file:Suppress("MaxLineLength") +@file:Suppress("MaxLineLength", "LargeClass") package com.wire.android.ui.authentication.login.email @@ -44,6 +44,8 @@ import com.wire.android.util.ui.CountdownTimer import com.wire.kalium.common.error.CoreFailure import com.wire.kalium.common.error.NetworkFailure import com.wire.kalium.logic.CoreLogic +import com.wire.kalium.logic.PrepareUserSessionResult +import com.wire.kalium.logic.UserSessionPreparationFailure import com.wire.kalium.logic.configuration.server.CommonApiVersionType import com.wire.kalium.logic.configuration.server.ServerConfig import com.wire.kalium.logic.data.auth.AccountInfo @@ -69,6 +71,7 @@ import com.wire.kalium.logic.feature.auth.verification.RequestSecondFactorVerifi import com.wire.kalium.logic.feature.client.ClientScope import com.wire.kalium.logic.feature.client.GetOrRegisterClientUseCase import com.wire.kalium.logic.feature.client.RegisterClientResult +import com.wire.kalium.logic.feature.UserSessionScope import com.wire.kalium.logic.feature.session.CurrentSessionResult import com.wire.kalium.logic.feature.session.CurrentSessionUseCase import com.wire.kalium.logic.feature.session.DeleteSessionUseCase @@ -624,6 +627,30 @@ class LoginEmailViewModelTest { } } + @Test + fun `given new session preparation fails, when canceling login, then preserve database and restore previous session`() = runTest { + val newUserId = UserId("newUserId", "domain") + val previousUserId = UserId("previousUserId", "domain") + val (arrangement, viewModel) = Arrangement() + .withPreparationFailure(UserSessionPreparationFailure.SupportRequired) + .withUpdateCurrentSessionReturning(UpdateCurrentSessionUseCase.Result.Success) + .arrange() + viewModel.loginJobData.value = LoginJobData( + job = mockk(relaxUnitFun = true), + previousSessionUserId = previousUserId, + newSessionUserId = newUserId, + ) + + viewModel.cancelLogin() + advanceUntilIdle() + + coVerify(exactly = 0) { + arrangement.logoutUseCase(any(), any()) + arrangement.deleteSessionUseCase(any()) + } + coVerify(exactly = 1) { arrangement.updateCurrentSessionUseCase(previousUserId) } + } + @Test fun `given no previous session, when canceling login, then update current session to null`() = runTest { // given @@ -831,6 +858,9 @@ class LoginEmailViewModelTest { @MockK internal lateinit var coreLogic: CoreLogic + @MockK + internal lateinit var userSessionScope: UserSessionScope + @MockK internal lateinit var requestSecondFactorCodeUseCase: RequestSecondFactorVerificationCodeUseCase @@ -862,6 +892,9 @@ class LoginEmailViewModelTest { every { qualifiedIdMapper.fromStringToQualifiedID(any()) } returns USER_ID every { savedInputStore.userIdentifier = any() } returns Unit every { coreLogic.getGlobalScope().validateEmailUseCase } returns validateEmailUseCase + coEvery { coreLogic.prepareUserSession(any()) } returns preparationSuccess(userSessionScope) + every { userSessionScope.users } returns userScope + every { userSessionScope.client } returns clientScope every { coreLogic.getSessionScope(any()).users } returns userScope every { userScope.persistSelfUserEmail } returns persistSelfUserEmailUseCase every { clientScopeProviderFactory.create(any()).clientScope } returns clientScope @@ -870,6 +903,7 @@ class LoginEmailViewModelTest { every { authenticationScope.login } returns loginUseCase every { authenticationScope.requestSecondFactorVerificationCode } returns requestSecondFactorCodeUseCase every { coreLogic.versionedAuthenticationScope(any()) } returns autoVersionAuthScopeUseCase + every { userSessionScope.logout } returns logoutUseCase every { coreLogic.getSessionScope(any()).logout } returns logoutUseCase every { coreLogic.getGlobalScope().deleteSession } returns deleteSessionUseCase every { coreLogic.getGlobalScope().session.updateCurrentSession } returns updateCurrentSessionUseCase @@ -878,6 +912,11 @@ class LoginEmailViewModelTest { coEvery { countdownTimer.start(any(), any(), any()) } returns Unit } + private fun preparationSuccess(sessionScope: UserSessionScope): PrepareUserSessionResult.Success = + mockk().also { result -> + every { result.sessionScope } returns sessionScope + } + fun arrange() = this to LoginEmailViewModel( LoginNavArgs(loginPasswordPath = LoginPasswordPath(newServerConfig(1).links)), addAuthenticatedUserUseCase, @@ -939,6 +978,12 @@ class LoginEmailViewModelTest { } returns result } + fun withPreparationFailure(reason: UserSessionPreparationFailure) = apply { + val result = mockk() + every { result.reason } returns reason + coEvery { coreLogic.prepareUserSession(any()) } returns result + } + fun withDeleteSessionReturning(result: DeleteSessionUseCase.Result) = apply { coEvery { deleteSessionUseCase(any()) diff --git a/app/src/test/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtensionTest.kt b/app/src/test/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtensionTest.kt index a91d0d29ae3..ed440b59f1b 100644 --- a/app/src/test/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtensionTest.kt +++ b/app/src/test/kotlin/com/wire/android/ui/authentication/login/sso/LoginSSOViewModelExtensionTest.kt @@ -21,12 +21,14 @@ package com.wire.android.ui.authentication.login.sso import com.wire.kalium.common.error.CoreFailure import com.wire.kalium.common.error.StorageFailure import com.wire.kalium.logic.CoreLogic +import com.wire.kalium.logic.PrepareUserSessionResult import com.wire.kalium.logic.data.auth.AccountTokens import com.wire.kalium.logic.data.logout.LogoutReason import com.wire.kalium.logic.data.session.StoreSessionParam import com.wire.kalium.logic.data.user.UserId import com.wire.kalium.logic.feature.auth.AddAuthenticatedUserUseCase import com.wire.kalium.logic.feature.auth.LogoutUseCase +import com.wire.kalium.logic.feature.UserSessionScope import com.wire.kalium.logic.feature.session.DeleteSessionUseCase import io.mockk.coEvery import io.mockk.coVerifyOrder @@ -117,8 +119,11 @@ class LoginSSOViewModelExtensionTest { val coreLogic = mockk() val logout = mockk() val deleteSession = mockk() + val userSessionScope = mockk() init { + coEvery { coreLogic.prepareUserSession(userId) } returns preparationSuccess(userSessionScope) + every { userSessionScope.logout } returns logout every { coreLogic.getSessionScope(userId).logout } returns logout every { coreLogic.getGlobalScope().deleteSession } returns deleteSession coEvery { logout(LogoutReason.SELF_HARD_LOGOUT, true) } returns Unit @@ -137,5 +142,10 @@ class LoginSSOViewModelExtensionTest { } fun arrange() = this to LoginSSOViewModelExtension(addAuthenticatedUser, coreLogic, false) + + private fun preparationSuccess(sessionScope: UserSessionScope): PrepareUserSessionResult.Success = + mockk().also { result -> + every { result.sessionScope } returns sessionScope + } } } diff --git a/app/src/test/kotlin/com/wire/android/ui/newauthentication/login/NewLoginViewModelTest.kt b/app/src/test/kotlin/com/wire/android/ui/newauthentication/login/NewLoginViewModelTest.kt index 40a3d1dc20c..71d2124af2a 100644 --- a/app/src/test/kotlin/com/wire/android/ui/newauthentication/login/NewLoginViewModelTest.kt +++ b/app/src/test/kotlin/com/wire/android/ui/newauthentication/login/NewLoginViewModelTest.kt @@ -840,7 +840,7 @@ class NewLoginViewModelTest { fun withRegisterClientReturning(result: RegisterClientResult) = apply { coEvery { - loginViewModelExtension.registerClient(any(), any(), any(), any()) + loginViewModelExtension.registerClient(any(), any(), any(), any()) } returns result }