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") })
+ }
+ }
+}