From 635c99751672d79e8e27bfc212217dc2d6f71a64 Mon Sep 17 00:00:00 2001 From: Greyson Parrelli Date: Mon, 31 Aug 2026 21:18:27 +0000 Subject: [PATCH] Detect re-used last-resort kyber key sets during decryption instead of at flush time. --- .../securesms/database/KyberPreKeyTable.kt | 28 +++++++++++---- .../database/LastResortKeyTupleTable.kt | 18 ++++++++++ .../protocol/BufferedKyberPreKeyStore.kt | 5 ++- .../database/KyberPreKeyTableTest.kt | 36 +++++++++++++++++++ .../protocol/BufferedKyberPreKeyStoreTest.kt | 19 ++++++++++ 5 files changed, 99 insertions(+), 7 deletions(-) diff --git a/app/src/main/java/org/thoughtcrime/securesms/database/KyberPreKeyTable.kt b/app/src/main/java/org/thoughtcrime/securesms/database/KyberPreKeyTable.kt index 16f3048211..4b856efe64 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/database/KyberPreKeyTable.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/database/KyberPreKeyTable.kt @@ -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") diff --git a/app/src/main/java/org/thoughtcrime/securesms/database/LastResortKeyTupleTable.kt b/app/src/main/java/org/thoughtcrime/securesms/database/LastResortKeyTupleTable.kt index ca908f8396..10d19dc82c 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/database/LastResortKeyTupleTable.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/database/LastResortKeyTupleTable.kt @@ -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) } + } } diff --git a/app/src/main/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStore.kt b/app/src/main/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStore.kt index a2fc43a65c..6e97f1ead2 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStore.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStore.kt @@ -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() } diff --git a/app/src/test/java/org/thoughtcrime/securesms/database/KyberPreKeyTableTest.kt b/app/src/test/java/org/thoughtcrime/securesms/database/KyberPreKeyTableTest.kt index fb52807636..7e53cd4c9d 100644 --- a/app/src/test/java/org/thoughtcrime/securesms/database/KyberPreKeyTableTest.kt +++ b/app/src/test/java/org/thoughtcrime/securesms/database/KyberPreKeyTableTest.kt @@ -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( diff --git a/app/src/test/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStoreTest.kt b/app/src/test/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStoreTest.kt index 2d5029e78c..9acc4f2b6f 100644 --- a/app/src/test/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStoreTest.kt +++ b/app/src/test/java/org/thoughtcrime/securesms/messages/protocol/BufferedKyberPreKeyStoreTest.kt @@ -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)