Detect re-used last-resort kyber key sets during decryption instead of at flush time.

This commit is contained in:
Greyson Parrelli
2026-09-30 13:34:01 -03:00
committed by Alex Hart
parent c1d43b31e6
commit 635c997516
5 changed files with 99 additions and 7 deletions
@@ -127,12 +127,7 @@ class KyberPreKeyTable(context: Context, databaseHelper: SignalDatabase) : Datab
*/
fun handleMarkKyberPreKeyUsed(serviceId: ServiceId, kyberPreKeyId: Int, signedPreKeyId: Int, baseKey: ECPublicKey) {
writableDatabase.withinTransaction { db ->
val lastResortRowId = db
.select(ID)
.from(TABLE_NAME)
.where("$ACCOUNT_ID = ? AND $KEY_ID = ? AND $LAST_RESORT = ?", serviceId.toAccountId(), kyberPreKeyId, 1)
.run()
.readToSingleInt(-1)
val lastResortRowId = getLastResortRowId(db, serviceId, kyberPreKeyId)
if (lastResortRowId < 0) {
db.delete("$TABLE_NAME INDEXED BY $INDEX_ACCOUNT_KEY")
@@ -148,6 +143,27 @@ class KyberPreKeyTable(context: Context, databaseHelper: SignalDatabase) : Datab
}
}
/**
* Whether we've already marked the given last-resort key set as used. If we have, then the sender is re-using a base key, and the message
* should be rejected rather than decrypted.
*
* Always false for non-last-resort keys, since those are simply deleted when used.
*/
fun hasUsedLastResortKeySet(serviceId: ServiceId, kyberPreKeyId: Int, signedPreKeyId: Int, baseKey: ECPublicKey): Boolean {
val lastResortRowId = getLastResortRowId(readableDatabase, serviceId, kyberPreKeyId)
return lastResortRowId >= 0 && SignalDatabase.lastResortKeyTuples.exists(lastResortRowId, signedPreKeyId, baseKey)
}
private fun getLastResortRowId(db: SQLiteDatabase, serviceId: ServiceId, kyberPreKeyId: Int): Int {
return db
.select(ID)
.from(TABLE_NAME)
.where("$ACCOUNT_ID = ? AND $KEY_ID = ? AND $LAST_RESORT = ?", serviceId.toAccountId(), kyberPreKeyId, 1)
.run()
.readToSingleInt(-1)
}
fun delete(serviceId: ServiceId, keyId: Int) {
writableDatabase
.delete("$TABLE_NAME INDEXED BY $INDEX_ACCOUNT_KEY")
@@ -9,6 +9,9 @@ import android.content.Context
import android.database.sqlite.SQLiteConstraintException
import org.signal.core.util.insertInto
import org.signal.core.util.logging.Log
import org.signal.core.util.readToList
import org.signal.core.util.requireNonNullBlob
import org.signal.core.util.select
import org.signal.libsignal.protocol.ReusedBaseKeyException
import org.signal.libsignal.protocol.ecc.ECPublicKey
@@ -59,4 +62,19 @@ class LastResortKeyTupleTable(context: Context, databaseHelper: SignalDatabase)
throw ReusedBaseKeyException(e)
}
}
/**
* Whether we've already recorded the given Last-resort tuple.
*/
fun exists(kyberPreKeyRowId: Int, signedKeyId: Int, publicKey: ECPublicKey): Boolean {
val serialized = publicKey.serialize()
return readableDatabase
.select(PUBLIC_KEY)
.from(TABLE_NAME)
.where("$KYBER_PREKEY = ? AND $SIGNED_KEY_ID = ?", kyberPreKeyRowId, signedKeyId)
.run()
.readToList { it.requireNonNullBlob(PUBLIC_KEY) }
.any { it.contentEquals(serialized) }
}
}
@@ -79,7 +79,10 @@ class BufferedKyberPreKeyStore(private val selfServiceId: ServiceId) : SignalSer
store.remove(kyberPreKeyId)
removedIfNotLastResort += Triple(kyberPreKeyId, signedPreKeyId, publicKey)
} else {
if (!lastResortKeyTuples.add(Triple(kyberPreKeyId, signedPreKeyId, publicKey))) {
// We don't have all the tuples in memory, so finding conflicts requires going to disk.
if (!lastResortKeyTuples.add(Triple(kyberPreKeyId, signedPreKeyId, publicKey)) ||
SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(selfServiceId, kyberPreKeyId, signedPreKeyId, publicKey)
) {
throw ReusedBaseKeyException()
}
@@ -7,8 +7,10 @@ package org.thoughtcrime.securesms.database
import android.app.Application
import org.junit.Assert.assertEquals
import org.junit.Assert.assertFalse
import org.junit.Assert.assertNotNull
import org.junit.Assert.assertNull
import org.junit.Assert.assertTrue
import org.junit.Rule
import org.junit.Test
import org.junit.runner.RunWith
@@ -200,6 +202,40 @@ class KyberPreKeyTableTest {
)
}
@Test
fun hasUsedLastResortKeySet_trueOnlyForKeySetsWeveAlreadySeen() {
insertTestRecord(aci, id = 1, staleTime = 10, lastResort = true)
val publicKey = generateECPublicKey()
assertFalse(SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(aci, kyberPreKeyId = 1, signedPreKeyId = 1, baseKey = publicKey))
SignalDatabase.kyberPreKeys.handleMarkKyberPreKeyUsed(
serviceId = aci,
kyberPreKeyId = 1,
signedPreKeyId = 1,
baseKey = publicKey
)
assertTrue(SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(aci, kyberPreKeyId = 1, signedPreKeyId = 1, baseKey = publicKey))
assertFalse(SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(aci, kyberPreKeyId = 1, signedPreKeyId = 2, baseKey = publicKey))
assertFalse(SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(aci, kyberPreKeyId = 1, signedPreKeyId = 1, baseKey = generateECPublicKey()))
}
@Test
fun hasUsedLastResortKeySet_falseForNonLastResortKeys() {
insertTestRecord(aci, id = 1, staleTime = 10, lastResort = false)
val publicKey = generateECPublicKey()
SignalDatabase.kyberPreKeys.handleMarkKyberPreKeyUsed(
serviceId = aci,
kyberPreKeyId = 1,
signedPreKeyId = 1,
baseKey = publicKey
)
assertFalse(SignalDatabase.kyberPreKeys.hasUsedLastResortKeySet(aci, kyberPreKeyId = 1, signedPreKeyId = 1, baseKey = publicKey))
}
private fun insertTestRecord(account: ServiceId, id: Int, staleTime: Long = 0, lastResort: Boolean = false) {
val kemKeyPair = KEMKeyPair.generate(KEMKeyType.KYBER_1024)
SignalDatabase.kyberPreKeys.insert(
@@ -70,6 +70,25 @@ class BufferedKyberPreKeyStoreTest {
)
}
@Test(expected = ReusedBaseKeyException::class)
fun givenALastResortKeyUsedInAnEarlierBatch_whenIMarkKyberPreKeyUsed_thenIExpectException() {
insertLastResortKey(id = 1)
val publicKey = generateECPublicKey()
SignalDatabase.kyberPreKeys.handleMarkKyberPreKeyUsed(
serviceId = aci,
kyberPreKeyId = 1,
signedPreKeyId = 2,
baseKey = publicKey
)
testSubject.markKyberPreKeyUsed(
kyberPreKeyId = 1,
signedPreKeyId = 2,
publicKey = publicKey
)
}
@Test
fun givenAMarkedLastResortKey_whenIFlushTwice_thenIExpectOnlyOneWrite() {
insertLastResortKey(id = 1)