From d72b8b730356c3dcf6f232b00eae4ae8b0952b47 Mon Sep 17 00:00:00 2001 From: Marten Rebane Date: Thu, 10 Sep 2026 11:42:33 +0300 Subject: [PATCH] Fix reusing previous Web eID request state --- .../DigiDoc/viewmodel/WebEidViewModelTest.kt | 122 ++++++++++++++++++ .../DigiDoc/domain/preferences/DataStore.kt | 2 + .../DigiDoc/fragment/screen/WebEidScreen.kt | 4 + .../ria/DigiDoc/viewmodel/WebEidViewModel.kt | 57 ++++++-- 4 files changed, 171 insertions(+), 14 deletions(-) diff --git a/app/src/androidTest/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModelTest.kt b/app/src/androidTest/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModelTest.kt index 6af3d4547..e0ea2d000 100644 --- a/app/src/androidTest/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModelTest.kt +++ b/app/src/androidTest/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModelTest.kt @@ -31,11 +31,15 @@ import ee.ria.DigiDoc.webEid.WebEidSignService import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.async import kotlinx.coroutines.flow.first +import kotlinx.coroutines.flow.toList +import kotlinx.coroutines.launch import kotlinx.coroutines.test.UnconfinedTestDispatcher +import kotlinx.coroutines.test.advanceUntilIdle import kotlinx.coroutines.test.runTest import org.json.JSONObject import org.junit.Assert.assertEquals import org.junit.Assert.assertNotNull +import org.junit.Assert.assertNull import org.junit.Before import org.junit.Rule import org.junit.Test @@ -79,12 +83,130 @@ class WebEidViewModelTest { private val signingCertBase64 = signingCertBase64Raw.replace("\\s+".toRegex(), "") + private val validAuthUri = + "web-eid-mobile://auth#eyJjaGFsbGVuZ2UiOiJ0ZXN0LWNoYWxsZW5nZS0wMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMCIsImxvZ2luVXJpIjoiaHR0cHM6Ly9leGFtcGxlLmNvbS9yZXNwb25zZSIsImdldFNpZ25pbmdDZXJ0aWZpY2F0ZSI6dHJ1ZX0" + + private val validCertUri = "web-eid-mobile://cert#eyJyZXNwb25zZVVyaSI6Imh0dHBzOi8vZXhhbXBsZS5jb20vcmVzcG9uc2UifQ" + @Before fun setup() { MockitoAnnotations.openMocks(this) viewModel = WebEidViewModel(authService, signService) } + @Test + fun webEidViewModel_handleCertificate_clearsPreviousSignRequest() { + runTest { + viewModel.handleSign(Uri.parse(createSignUri(signingCertBase64))) + assertNotNull(viewModel.signRequest.value) + + viewModel.handleCertificate(Uri.parse(validCertUri)) + + assertNull(viewModel.signRequest.value) + assertNotNull(viewModel.certificateRequest.value) + } + } + + @Test + fun webEidViewModel_handleSign_clearsPreviousCertificateRequest() { + runTest { + viewModel.handleCertificate(Uri.parse(validCertUri)) + assertNotNull(viewModel.certificateRequest.value) + + viewModel.handleSign(Uri.parse(createSignUri(signingCertBase64))) + + assertNull(viewModel.certificateRequest.value) + assertNotNull(viewModel.signRequest.value) + } + } + + @Test + fun webEidViewModel_handleAuth_clearsPreviousSignRequest() { + runTest { + viewModel.handleSign(Uri.parse(createSignUri(signingCertBase64))) + assertNotNull(viewModel.signRequest.value) + + viewModel.handleAuth(Uri.parse(validAuthUri)) + + assertNull(viewModel.signRequest.value) + assertNotNull(viewModel.authRequest.value) + } + } + + @Test + fun webEidViewModel_handleUnknown_clearsPreviousRequests() { + runTest { + viewModel.handleSign(Uri.parse(createSignUri(signingCertBase64))) + assertNotNull(viewModel.signRequest.value) + + viewModel.handleUnknown(Uri.parse("web-eid-mobile://unknown#dGVzdA")) + + assertNull(viewModel.signRequest.value) + assertNull(viewModel.certificateRequest.value) + assertNull(viewModel.authRequest.value) + } + } + + @Test + fun webEidViewModel_newRequest_clearsPreviousDialogError() { + runTest { + viewModel.handleUnknown(Uri.parse("web-eid-mobile://unknown#dGVzdA")) + assertEquals(R.string.web_eid_invalid_request_error, viewModel.dialogError.value) + + viewModel.handleAuth(Uri.parse(validAuthUri)) + + assertEquals(0, viewModel.dialogError.value) + } + } + + @OptIn(ExperimentalCoroutinesApi::class) + @Test + fun webEidViewModel_secondResponseForSameFlow_isIgnored() { + runTest(UnconfinedTestDispatcher()) { + val emitted = mutableListOf() + val job = launch { viewModel.relyingPartyResponseEvents.toList(emitted) } + + viewModel.handleAuth(Uri.parse(validAuthUri)) + viewModel.handleUserCancelled() + viewModel.handleUserCancelled() + + advanceUntilIdle() + job.cancel() + assertEquals(1, emitted.size) + } + } + + @OptIn(ExperimentalCoroutinesApi::class) + @Test + fun webEidViewModel_handleAuth_permitsOneResponsePerRequest() { + runTest(UnconfinedTestDispatcher()) { + val emitted = mutableListOf() + val job = launch { viewModel.relyingPartyResponseEvents.toList(emitted) } + + viewModel.handleAuth(Uri.parse(validAuthUri)) + viewModel.handleUserCancelled() + viewModel.handleAuth(Uri.parse(validAuthUri)) + viewModel.handleUserCancelled() + + advanceUntilIdle() + job.cancel() + assertEquals(2, emitted.size) + } + } + + @Test + fun webEidViewModel_response_isDeliveredWhenCollectedAfterItIsSent() { + runTest { + viewModel.handleAuth(Uri.parse(validAuthUri)) + viewModel.handleUserCancelled() + + val received = viewModel.relyingPartyResponseEvents.first() + + assertNotNull(received) + assert(received.toString().startsWith("https://example.com/response#")) + } + } + @Test fun webEidViewModel_handleAuth_parsesAuthUriAndSetsStateFlow() { runTest { diff --git a/app/src/main/kotlin/ee/ria/DigiDoc/domain/preferences/DataStore.kt b/app/src/main/kotlin/ee/ria/DigiDoc/domain/preferences/DataStore.kt index cb13b8417..7dc2b0dfe 100644 --- a/app/src/main/kotlin/ee/ria/DigiDoc/domain/preferences/DataStore.kt +++ b/app/src/main/kotlin/ee/ria/DigiDoc/domain/preferences/DataStore.kt @@ -130,6 +130,7 @@ class DataStore } fun clearTemporaryCanNumber() { + debugLog(logTag, "Clearing temporary CAN") runEncryptedWrite(context, "Unable to clear temporary CAN") { EncryptedPreferences.putString( context, @@ -153,6 +154,7 @@ class DataStore }.toBoolean() fun setWebEidSessionActive(active: Boolean) { + debugLog(logTag, "Setting Web eID session active: $active") runEncryptedWrite(context, "Unable to save Web eID session state") { EncryptedPreferences.putString(context, "web_eid_session_active", active.toString()) } diff --git a/app/src/main/kotlin/ee/ria/DigiDoc/fragment/screen/WebEidScreen.kt b/app/src/main/kotlin/ee/ria/DigiDoc/fragment/screen/WebEidScreen.kt index 9853b1f3a..927c587d2 100644 --- a/app/src/main/kotlin/ee/ria/DigiDoc/fragment/screen/WebEidScreen.kt +++ b/app/src/main/kotlin/ee/ria/DigiDoc/fragment/screen/WebEidScreen.kt @@ -219,10 +219,12 @@ fun WebEidScreen( rememberMe = rememberMe, isWebEidAuthenticating = isWebEidAuthenticating, onError = { + sharedSettingsViewModel.dataStore.setWebEidSessionActive(false) isWebEidAuthenticating = false cancelWebEidAuthenticateAction() }, onSuccess = { + sharedSettingsViewModel.dataStore.setWebEidSessionActive(false) isWebEidAuthenticating = false navController.navigateUp() }, @@ -277,6 +279,7 @@ fun WebEidScreen( showPinField = false, isWebEidAuthenticating = isWebEidAuthenticating, onError = { + sharedSettingsViewModel.dataStore.setWebEidSessionActive(false) isWebEidAuthenticating = false cancelWebEidSignAction() }, @@ -308,6 +311,7 @@ fun WebEidScreen( isWebEidAuthenticating = isWebEidAuthenticating, isCanNumberReadOnly = hasStoredCanNumber, onError = { + sharedSettingsViewModel.dataStore.setWebEidSessionActive(false) isWebEidAuthenticating = false cancelWebEidSignAction() }, diff --git a/app/src/main/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModel.kt b/app/src/main/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModel.kt index 8f55ac6d1..c58fde17c 100644 --- a/app/src/main/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModel.kt +++ b/app/src/main/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModel.kt @@ -25,6 +25,7 @@ import android.net.Uri import androidx.lifecycle.ViewModel import dagger.hilt.android.lifecycle.HiltViewModel import ee.ria.DigiDoc.R +import ee.ria.DigiDoc.utilsLib.logging.LoggingUtil.Companion.debugLog import ee.ria.DigiDoc.utilsLib.logging.LoggingUtil.Companion.errorLog import ee.ria.DigiDoc.webEid.WebEidAuthService import ee.ria.DigiDoc.webEid.WebEidSignService @@ -35,12 +36,12 @@ import ee.ria.DigiDoc.webEid.exception.WebEidErrorCode import ee.ria.DigiDoc.webEid.exception.WebEidException import ee.ria.DigiDoc.webEid.utils.WebEidRequestParser import ee.ria.DigiDoc.webEid.utils.WebEidResponseUtil -import kotlinx.coroutines.flow.MutableSharedFlow +import kotlinx.coroutines.channels.Channel +import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.MutableStateFlow -import kotlinx.coroutines.flow.SharedFlow import kotlinx.coroutines.flow.StateFlow -import kotlinx.coroutines.flow.asSharedFlow import kotlinx.coroutines.flow.asStateFlow +import kotlinx.coroutines.flow.receiveAsFlow import org.json.JSONObject import javax.inject.Inject @@ -59,19 +60,44 @@ class WebEidViewModel val certificateRequest: StateFlow = _certificateRequest.asStateFlow() private val _signRequest = MutableStateFlow(null) val signRequest: StateFlow = _signRequest.asStateFlow() - private val _relyingPartyResponseEvents = MutableSharedFlow() - val relyingPartyResponseEvents: SharedFlow = _relyingPartyResponseEvents.asSharedFlow() + private val _relyingPartyResponseEvents = Channel(Channel.BUFFERED) + val relyingPartyResponseEvents: Flow = _relyingPartyResponseEvents.receiveAsFlow() private val _dialogError = MutableStateFlow(0) val dialogError: StateFlow = _dialogError + private var hasRespondedToRelyingParty = false + + private fun resetRequests() { + debugLog( + logTag, + "Resetting Web eID request state. Had auth request: ${_authRequest.value != null}. " + + "Had certificate request: ${_certificateRequest.value != null}. " + + "Had sign request: ${_signRequest.value != null}", + ) + _authRequest.value = null + _certificateRequest.value = null + _signRequest.value = null + _dialogError.value = 0 + hasRespondedToRelyingParty = false + } + + private suspend fun sendResponse(responseUri: Uri) { + if (hasRespondedToRelyingParty) { + errorLog(logTag, "Ignoring duplicate response to relying party") + return + } + hasRespondedToRelyingParty = true + _relyingPartyResponseEvents.send(responseUri) + } suspend fun handleAuth(uri: Uri) { + resetRequests() try { _authRequest.value = WebEidRequestParser.parseAuthUri(uri) } catch (e: WebEidException) { errorLog(logTag, "Invalid Web eID authentication request: $uri", e) val errorPayload = WebEidResponseUtil.createErrorPayload(e.errorCode, e.message) val responseUri = WebEidResponseUtil.createResponseUri(e.responseUri, errorPayload) - _relyingPartyResponseEvents.emit(responseUri) + sendResponse(responseUri) } catch (e: Exception) { errorLog(logTag, "Unable parse Web eID authentication request: $uri", e) _dialogError.value = R.string.web_eid_invalid_auth_request_error @@ -79,6 +105,7 @@ class WebEidViewModel } fun handleCertificate(uri: Uri) { + resetRequests() try { _certificateRequest.value = WebEidRequestParser.parseCertificateUri(uri) } catch (e: Exception) { @@ -88,13 +115,14 @@ class WebEidViewModel } suspend fun handleSign(uri: Uri) { + resetRequests() try { _signRequest.value = WebEidRequestParser.parseSignUri(uri) } catch (e: WebEidException) { errorLog(logTag, "Invalid Web eID signing request: $uri", e) val errorPayload = WebEidResponseUtil.createErrorPayload(e.errorCode, e.message) val responseUri = WebEidResponseUtil.createResponseUri(e.responseUri, errorPayload) - _relyingPartyResponseEvents.emit(responseUri) + sendResponse(responseUri) } catch (e: Exception) { errorLog(logTag, "Unable parse Web eID signing request: $uri", e) _dialogError.value = R.string.web_eid_invalid_request_error @@ -102,6 +130,7 @@ class WebEidViewModel } fun handleUnknown(uri: Uri) { + resetRequests() errorLog(logTag, "Unable parse Web eID request: $uri") _dialogError.value = R.string.web_eid_invalid_request_error } @@ -123,7 +152,7 @@ class WebEidViewModel ) val payload = JSONObject().put("authToken", token) val responseUri = WebEidResponseUtil.createResponseUri(loginUri, payload) - _relyingPartyResponseEvents.emit(responseUri) + sendResponse(responseUri) } catch (e: Exception) { errorLog(logTag, "Unexpected error building auth token", e) val errorPayload = @@ -132,7 +161,7 @@ class WebEidViewModel "Unexpected error", ) val responseUri = WebEidResponseUtil.createResponseUri(loginUri, errorPayload) - _relyingPartyResponseEvents.emit(responseUri) + sendResponse(responseUri) } } @@ -147,7 +176,7 @@ class WebEidViewModel try { val payload = signService.buildCertificatePayload(signingCert) val response = WebEidResponseUtil.createResponseUri(responseUri, payload) - _relyingPartyResponseEvents.emit(response) + sendResponse(response) } catch (e: Exception) { errorLog(logTag, "Unexpected error building certificate payload", e) val errorPayload = @@ -156,7 +185,7 @@ class WebEidViewModel "Unexpected error", ) val errorUri = WebEidResponseUtil.createResponseUri(responseUri, errorPayload) - _relyingPartyResponseEvents.emit(errorUri) + sendResponse(errorUri) } } @@ -177,7 +206,7 @@ class WebEidViewModel hashFunction, ) val response = WebEidResponseUtil.createResponseUri(responseUri, payload) - _relyingPartyResponseEvents.emit(response) + sendResponse(response) } catch (e: Exception) { errorLog(logTag, "Unexpected error building sign payload", e) val errorPayload = @@ -186,7 +215,7 @@ class WebEidViewModel "Unexpected error", ) val errorUri = WebEidResponseUtil.createResponseUri(responseUri, errorPayload) - _relyingPartyResponseEvents.emit(errorUri) + sendResponse(errorUri) } } @@ -211,7 +240,7 @@ class WebEidViewModel val errorUri = WebEidResponseUtil.createResponseUri(responseUri, errorPayload) - _relyingPartyResponseEvents.emit(errorUri) + sendResponse(errorUri) } catch (e: Exception) { errorLog(logTag, "Failed to send cancel response", e) }