diff --git a/key/data/src/main/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImpl.kt b/key/data/src/main/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImpl.kt index 78cfc30c0..6cc54e42a 100644 --- a/key/data/src/main/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImpl.kt +++ b/key/data/src/main/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImpl.kt @@ -73,6 +73,7 @@ class PublicAddressRepositoryImpl @Inject constructor( private suspend fun insertOrUpdate(publicAddress: PublicAddress) = db.inTransaction { publicAddressDao.insertOrUpdate(publicAddress.toEntity()) + publicAddressKeyDao.deleteByEmail(publicAddress.email) publicAddressKeyDao.insertOrUpdate(*publicAddress.keys.toEntityList().toTypedArray()) } diff --git a/key/data/src/test/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImplTest.kt b/key/data/src/test/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImplTest.kt new file mode 100644 index 000000000..c702a645f --- /dev/null +++ b/key/data/src/test/kotlin/me/proton/core/key/data/repository/PublicAddressRepositoryImplTest.kt @@ -0,0 +1,151 @@ +/* + * Copyright (c) 2022 Proton Technologies 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 . + */ + +package me.proton.core.key.data.repository + +import io.mockk.coEvery +import io.mockk.coJustRun +import io.mockk.coVerify +import io.mockk.every +import io.mockk.mockk +import io.mockk.slot +import kotlinx.coroutines.flow.flowOf +import kotlinx.coroutines.test.runTest +import me.proton.core.domain.entity.UserId +import me.proton.core.key.data.api.KeyApi +import me.proton.core.key.data.db.PublicAddressDao +import me.proton.core.key.data.db.PublicAddressDatabase +import me.proton.core.key.data.db.PublicAddressKeyDao +import me.proton.core.key.data.db.PublicAddressWithKeysDao +import me.proton.core.key.data.entity.PublicAddressKeyEntity +import me.proton.core.network.data.ApiManagerFactory +import me.proton.core.network.data.ApiProvider +import me.proton.core.network.domain.ApiManager +import me.proton.core.network.domain.ApiResult +import me.proton.core.network.domain.session.SessionId +import me.proton.core.network.domain.session.SessionProvider +import me.proton.core.test.kotlin.TestDispatcherProvider +import org.junit.Assert.assertEquals +import org.junit.Before +import org.junit.Test + +class PublicAddressRepositoryImplTest { + + private lateinit var repositoryImpl: PublicAddressRepositoryImpl + private val db: PublicAddressDatabase = mockk() + private val sessionProvider = mockk() + private val keyApi = mockk() + private val apiFactory = mockk() + private lateinit var apiProvider: ApiProvider + + private val testSessionId = SessionId("test-session-id") + private val testUserId = UserId("test-user-id") + + private val dispatcherProvider = TestDispatcherProvider + + private val publicAddressDao = mockk() + private val publicAddressWithKeysDao = mockk() + private val publicAddressKeyDao = mockk() + + @Before + fun setUp() { + val apiManager = object : ApiManager { + override suspend fun invoke( + forceNoRetryOnConnectionErrors: Boolean, + block: suspend KeyApi.() -> T + ): ApiResult = ApiResult.Success(block.invoke(keyApi)) + } + + coEvery { sessionProvider.getSessionId(testUserId) } returns testSessionId + apiProvider = ApiProvider(apiFactory, sessionProvider, dispatcherProvider) + every { apiFactory.create(testSessionId, interfaceClass = KeyApi::class) } returns apiManager + coEvery { db.publicAddressDao() } returns publicAddressDao + coEvery { db.publicAddressKeyDao() } returns publicAddressKeyDao + coEvery { db.publicAddressWithKeysDao() } returns publicAddressWithKeysDao + val transactionLambda = slot Unit>() + coEvery { db.inTransaction(capture(transactionLambda)) } coAnswers { + transactionLambda.captured.invoke() + } + repositoryImpl = PublicAddressRepositoryImpl( + db, + apiProvider + ) + } + + @Test + fun `Old public keys are removed from local db`() = runTest { + // given + val testEmail = "email" + val storedKeys = mutableListOf("key1", "key2") + coJustRun { publicAddressDao.insertOrUpdate(any()) } + coEvery { publicAddressKeyDao.deleteByEmail(testEmail) } coAnswers { + storedKeys.clear() + } + val insertKeys = mutableListOf() + coEvery { publicAddressKeyDao.insertOrUpdate(*varargAll { insertKeys.add(it) }) } coAnswers { + insertKeys.forEach { key -> + if (key.email == testEmail) { + storedKeys.add(key.publicKey) + } + } + insertKeys.clear() + } + coEvery { + keyApi.getPublicAddressKeys(testEmail, any()) + } returns mockk { + every { toPublicAddress(testEmail) } returns mockk(relaxed = true) { + every { email } returns testEmail + every { keys } returns listOf( + mockk(relaxed = true) { + every { email } returns testEmail + every { publicKey.key } returns "key2" + }, + mockk(relaxed = true) { + every { email } returns testEmail + every { publicKey.key } returns "key3" + } + ) + } + } + coEvery { publicAddressWithKeysDao.findWithKeysByEmail(testEmail) } answers { + flowOf( + mockk(relaxed = true) { + every { entity.email } returns testEmail + every { keys } returns storedKeys.map { key -> + mockk(relaxed = true) { + every { publicKey } returns key + } + } + } + ) + } + val expectedKeys = listOf("key2", "key3") + // when + val address = repositoryImpl.getPublicAddress(testUserId, testEmail) + // then + assertEquals(expectedKeys, storedKeys) + assertEquals(expectedKeys, address.keys.map { it.publicKey.key }) + assertEquals(testEmail, address.email) + coVerify { + keyApi.getPublicAddressKeys(testEmail, any()) + publicAddressDao.insertOrUpdate(match { it.email == testEmail }) + publicAddressKeyDao.deleteByEmail(testEmail) + publicAddressKeyDao.insertOrUpdate(*varargAll { it.publicKey in listOf("key2", "key3") }) + } + } +}