Move BufferedKyberPreKeyStoreTest to a unit test.

This commit is contained in:
Greyson Parrelli
2026-09-30 13:34:01 -03:00
committed by Alex Hart
parent 38c62bce46
commit c1d43b31e6
3 changed files with 108 additions and 150 deletions
@@ -1,79 +0,0 @@
/*
* Copyright 2025 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.thoughtcrime.securesms.messages.protocol
import org.junit.Before
import org.junit.Rule
import org.junit.Test
import org.signal.core.models.ServiceId
import org.signal.libsignal.protocol.ReusedBaseKeyException
import org.thoughtcrime.securesms.keyvalue.SignalStore
import org.thoughtcrime.securesms.testing.SignalDatabaseRule
import org.thoughtcrime.securesms.util.KyberPreKeysTestUtil
class BufferedKyberPreKeyStoreTest {
@get:Rule
val harness = SignalDatabaseRule()
private lateinit var aci: ServiceId
private lateinit var testSubject: BufferedKyberPreKeyStore
private lateinit var dataStore: BufferedSignalServiceAccountDataStore
@Before
fun setUp() {
SignalStore.account.generateAciIdentityKeyIfNecessary()
aci = harness.localAci
testSubject = BufferedKyberPreKeyStore(aci)
dataStore = BufferedSignalServiceAccountDataStore(aci)
}
@Test
fun givenALastResortKey_whenIMarkKyberPreKeyUsed_thenIExpectNoIssues() {
KyberPreKeysTestUtil.insertTestRecord(aci, 1, lastResort = true)
val publicKey = KyberPreKeysTestUtil.generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
}
@Test(expected = ReusedBaseKeyException::class)
fun givenALastResortKey_whenIMarkKyberPreKeyUsedTwice_thenIExpectException() {
KyberPreKeysTestUtil.insertTestRecord(aci, 1, lastResort = true)
val publicKey = KyberPreKeysTestUtil.generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
}
@Test
fun givenAMarkedLastResortKey_whenIFlushTwice_thenIExpectNoIssues() {
KyberPreKeysTestUtil.insertTestRecord(aci, 1, lastResort = true)
val publicKey = KyberPreKeysTestUtil.generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
testSubject.flushToDisk(dataStore)
testSubject.flushToDisk(dataStore)
}
}
@@ -1,71 +0,0 @@
/*
* Copyright 2025 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.thoughtcrime.securesms.util
import org.junit.Assert.assertEquals
import org.signal.core.models.ServiceId
import org.signal.core.models.ServiceId.ACI
import org.signal.core.models.ServiceId.PNI
import org.signal.core.util.readToSingleObject
import org.signal.core.util.requireLongOrNull
import org.signal.core.util.select
import org.signal.core.util.update
import org.signal.libsignal.protocol.ecc.ECKeyPair
import org.signal.libsignal.protocol.ecc.ECPublicKey
import org.signal.libsignal.protocol.kem.KEMKeyPair
import org.signal.libsignal.protocol.kem.KEMKeyType
import org.signal.libsignal.protocol.state.KyberPreKeyRecord
import org.thoughtcrime.securesms.database.KyberPreKeyTable
import org.thoughtcrime.securesms.database.SignalDatabase
import java.security.SecureRandom
object KyberPreKeysTestUtil {
fun insertTestRecord(account: ServiceId, id: Int, staleTime: Long = 0, lastResort: Boolean = false) {
val kemKeyPair = KEMKeyPair.generate(KEMKeyType.KYBER_1024)
SignalDatabase.kyberPreKeys.insert(
serviceId = account,
keyId = id,
record = KyberPreKeyRecord(
id,
System.currentTimeMillis(),
kemKeyPair,
ECKeyPair.generate().privateKey.calculateSignature(kemKeyPair.publicKey.serialize())
),
lastResort = lastResort
)
val count = SignalDatabase.rawDatabase
.update(KyberPreKeyTable.TABLE_NAME)
.values(KyberPreKeyTable.STALE_TIMESTAMP to staleTime)
.where("${KyberPreKeyTable.ACCOUNT_ID} = ? AND ${KyberPreKeyTable.KEY_ID} = $id", account.toAccountId())
.run()
assertEquals(1, count)
}
fun getStaleTime(account: ServiceId, id: Int): Long? {
return SignalDatabase.rawDatabase
.select(KyberPreKeyTable.STALE_TIMESTAMP)
.from(KyberPreKeyTable.TABLE_NAME)
.where("${KyberPreKeyTable.ACCOUNT_ID} = ? AND ${KyberPreKeyTable.KEY_ID} = $id", account.toAccountId())
.run()
.readToSingleObject { it.requireLongOrNull(KyberPreKeyTable.STALE_TIMESTAMP) }
}
fun generateECPublicKey(): ECPublicKey {
val byteArray = ByteArray(ECPublicKey.KEY_SIZE - 1)
SecureRandom().nextBytes(byteArray)
return ECPublicKey.fromPublicKeyBytes(byteArray)
}
private fun ServiceId.toAccountId(): String {
return when (this) {
is ACI -> this.toString()
is PNI -> KyberPreKeyTable.PNI_ACCOUNT_ID
}
}
}
@@ -0,0 +1,108 @@
/*
* Copyright 2025 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.thoughtcrime.securesms.messages.protocol
import android.app.Application
import io.mockk.mockk
import io.mockk.verify
import org.junit.Rule
import org.junit.Test
import org.junit.runner.RunWith
import org.robolectric.RobolectricTestRunner
import org.robolectric.annotation.Config
import org.signal.core.models.ServiceId.ACI
import org.signal.libsignal.protocol.ReusedBaseKeyException
import org.signal.libsignal.protocol.ecc.ECKeyPair
import org.signal.libsignal.protocol.ecc.ECPublicKey
import org.signal.libsignal.protocol.kem.KEMKeyPair
import org.signal.libsignal.protocol.kem.KEMKeyType
import org.signal.libsignal.protocol.state.KyberPreKeyRecord
import org.thoughtcrime.securesms.database.SignalDatabase
import org.thoughtcrime.securesms.testutil.MockAppDependenciesRule
import org.thoughtcrime.securesms.testutil.SignalDatabaseRule
import org.whispersystems.signalservice.api.SignalServiceAccountDataStore
import java.util.UUID
@RunWith(RobolectricTestRunner::class)
@Config(manifest = Config.NONE, application = Application::class)
class BufferedKyberPreKeyStoreTest {
@get:Rule
val appDependencies = MockAppDependenciesRule()
@get:Rule
val signalDatabaseRule = SignalDatabaseRule()
private val aci: ACI = ACI.from(UUID.randomUUID())
private val testSubject = BufferedKyberPreKeyStore(aci)
private val dataStore: SignalServiceAccountDataStore = mockk(relaxed = true)
@Test
fun givenALastResortKey_whenIMarkKyberPreKeyUsed_thenIExpectNoIssues() {
insertLastResortKey(id = 1)
val publicKey = generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
}
@Test(expected = ReusedBaseKeyException::class)
fun givenALastResortKey_whenIMarkKyberPreKeyUsedTwice_thenIExpectException() {
insertLastResortKey(id = 1)
val publicKey = generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
}
@Test
fun givenAMarkedLastResortKey_whenIFlushTwice_thenIExpectOnlyOneWrite() {
insertLastResortKey(id = 1)
val publicKey = generateECPublicKey()
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
testSubject.flushToDisk(dataStore)
testSubject.flushToDisk(dataStore)
verify(exactly = 1) { dataStore.markKyberPreKeyUsed(1, 2, publicKey) }
}
private fun insertLastResortKey(id: Int) {
val kemKeyPair = KEMKeyPair.generate(KEMKeyType.KYBER_1024)
SignalDatabase.kyberPreKeys.insert(
serviceId = aci,
keyId = id,
record = KyberPreKeyRecord(
id,
System.currentTimeMillis(),
kemKeyPair,
ECKeyPair.generate().privateKey.calculateSignature(kemKeyPair.publicKey.serialize())
),
lastResort = true
)
}
private fun generateECPublicKey(): ECPublicKey {
return ECKeyPair.generate().publicKey
}
}