Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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<Uri>()
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<Uri>()
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 {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,7 @@ class DataStore
}

fun clearTemporaryCanNumber() {
debugLog(logTag, "Clearing temporary CAN")
runEncryptedWrite(context, "Unable to clear temporary CAN") {
EncryptedPreferences.putString(
context,
Expand All @@ -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())
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
},
Expand Down Expand Up @@ -277,6 +279,7 @@ fun WebEidScreen(
showPinField = false,
isWebEidAuthenticating = isWebEidAuthenticating,
onError = {
sharedSettingsViewModel.dataStore.setWebEidSessionActive(false)
isWebEidAuthenticating = false
cancelWebEidSignAction()
},
Expand Down Expand Up @@ -308,6 +311,7 @@ fun WebEidScreen(
isWebEidAuthenticating = isWebEidAuthenticating,
isCanNumberReadOnly = hasStoredCanNumber,
onError = {
sharedSettingsViewModel.dataStore.setWebEidSessionActive(false)
isWebEidAuthenticating = false
cancelWebEidSignAction()
},
Expand Down
57 changes: 43 additions & 14 deletions app/src/main/kotlin/ee/ria/DigiDoc/viewmodel/WebEidViewModel.kt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -59,26 +60,52 @@ class WebEidViewModel
val certificateRequest: StateFlow<WebEidCertificateRequest?> = _certificateRequest.asStateFlow()
private val _signRequest = MutableStateFlow<WebEidSignRequest?>(null)
val signRequest: StateFlow<WebEidSignRequest?> = _signRequest.asStateFlow()
private val _relyingPartyResponseEvents = MutableSharedFlow<Uri>()
val relyingPartyResponseEvents: SharedFlow<Uri> = _relyingPartyResponseEvents.asSharedFlow()
private val _relyingPartyResponseEvents = Channel<Uri>(Channel.BUFFERED)
val relyingPartyResponseEvents: Flow<Uri> = _relyingPartyResponseEvents.receiveAsFlow()
private val _dialogError = MutableStateFlow(0)
val dialogError: StateFlow<Int> = _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
}
}

fun handleCertificate(uri: Uri) {
resetRequests()
try {
_certificateRequest.value = WebEidRequestParser.parseCertificateUri(uri)
} catch (e: Exception) {
Expand All @@ -88,20 +115,22 @@ 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
}
}

fun handleUnknown(uri: Uri) {
resetRequests()
errorLog(logTag, "Unable parse Web eID request: $uri")
_dialogError.value = R.string.web_eid_invalid_request_error
}
Expand All @@ -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 =
Expand All @@ -132,7 +161,7 @@ class WebEidViewModel
"Unexpected error",
)
val responseUri = WebEidResponseUtil.createResponseUri(loginUri, errorPayload)
_relyingPartyResponseEvents.emit(responseUri)
sendResponse(responseUri)
}
}

Expand All @@ -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 =
Expand All @@ -156,7 +185,7 @@ class WebEidViewModel
"Unexpected error",
)
val errorUri = WebEidResponseUtil.createResponseUri(responseUri, errorPayload)
_relyingPartyResponseEvents.emit(errorUri)
sendResponse(errorUri)
}
}

Expand All @@ -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 =
Expand All @@ -186,7 +215,7 @@ class WebEidViewModel
"Unexpected error",
)
val errorUri = WebEidResponseUtil.createResponseUri(responseUri, errorPayload)
_relyingPartyResponseEvents.emit(errorUri)
sendResponse(errorUri)
}
}

Expand All @@ -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)
}
Expand Down
Loading