From 4ccad176e192d72211739e8f98ef7409b0e24978 Mon Sep 17 00:00:00 2001 From: jeremyd <4072+jeremyd@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:51:45 -0700 Subject: [PATCH] feat(mls): keep skipped-generation secrets across a restore The out-of-order cache in SecretTree lived only in memory. A member that opened a later message, then restarted, could never open the earlier ones: the ratchet had moved past them and their keys were gone. SecretTree now caches the ratchet secret of each skipped generation instead of the derived key/nonce (the form ts-mls keeps in unusedGenerations), and exports/imports it. MlsGroupState v5 persists both ratchets' skipped secrets; v1-v4 blobs still decode with none. --- .../quartz/mls/group/MlsGroup.kt | 3 + .../quartz/mls/group/MlsGroupState.kt | 49 +++++++++- .../quartz/mls/schedule/SecretTree.kt | 54 ++++++++--- .../schedule/SecretTreeSkippedSecretsTest.kt | 58 ++++++++++++ .../MlsGroupStateSkippedGenerationsTest.kt | 93 +++++++++++++++++++ .../quartz/mls/MlsGroupStateTest.kt | 8 +- 6 files changed, 245 insertions(+), 20 deletions(-) create mode 100644 quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeSkippedSecretsTest.kt create mode 100644 quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateSkippedGenerationsTest.kt diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroup.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroup.kt index 782b3effe4..412515b562 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroup.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroup.kt @@ -290,6 +290,8 @@ class MlsGroup private constructor( // restart that forgot it would leave the leaver in the tree with // the group's keys and nobody holding the proposal to evict them. pendingProposals = pendingProposals.toList(), + skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(), + skippedHandshakeSecrets = secretTree.exportSkippedHandshakeSecrets(), ) } @@ -4039,6 +4041,7 @@ class MlsGroup private constructor( val tree = RatchetTree.decodeTls(TlsReader(state.treeBytes)) val secretTree = SecretTree(state.encryptionSecret, tree.leafCount) secretTree.importSenderStates(state.senderRatchetStates) + secretTree.importSkippedSecrets(state.skippedApplicationSecrets, state.skippedHandshakeSecrets) return MlsGroup( groupContext = state.groupContext, diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroupState.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroupState.kt index 4b4907cc2c..7ea4ca5086 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroupState.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/group/MlsGroupState.kt @@ -83,6 +83,17 @@ data class MlsGroupState( * the proposal that would evict them. */ val pendingProposals: List = emptyList(), + /** + * Secrets of skipped APPLICATION generations, (leafIndex, generation) -> + * secret (STATE_VERSION 5+). + * + * A message that arrives after a later one from the same sender needs + * its generation's secret, and the ratchet has already moved past it. + * Without these a restart between the two loses that message for good. + */ + val skippedApplicationSecrets: Map, ByteArray> = emptyMap(), + /** Same for the HANDSHAKE ratchet (STATE_VERSION 5+). */ + val skippedHandshakeSecrets: Map, ByteArray> = emptyMap(), ) { fun encodeTls(): ByteArray { val writer = TlsWriter() @@ -156,6 +167,10 @@ data class MlsGroupState( writer.putOpaqueVarInt(pending.authenticatedContentBytes ?: ByteArray(0)) } + // Skipped-generation secrets (STATE_VERSION 5+). + writeSkippedSecrets(writer, skippedApplicationSecrets) + writeSkippedSecrets(writer, skippedHandshakeSecrets) + return writer.toByteArray() } @@ -182,8 +197,11 @@ data class MlsGroupState( * v4: appends [pendingProposals] so a departing member's staged * `SelfRemove` survives a restart instead of leaving them in the * tree. Older blobs decode with an empty pool. + * v5: appends [skippedApplicationSecrets] and [skippedHandshakeSecrets] + * so an out-of-order message survives a restart. Older blobs + * decode with none, as before. */ - private const val STATE_VERSION = 4 + private const val STATE_VERSION = 5 fun decodeTls(data: ByteArray): MlsGroupState { val reader = TlsReader(data) @@ -284,6 +302,10 @@ data class MlsGroupState( emptyList() } + // v5+: skipped-generation secrets. Absent for older blobs. + val skippedApplicationSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap() + val skippedHandshakeSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap() + return MlsGroupState( groupContext = groupContext, treeBytes = treeBytes, @@ -297,8 +319,33 @@ data class MlsGroupState( senderRatchetStates = senderRatchetStates, pathPrivateKeys = pathPrivateKeys, pendingProposals = pendingProposals, + skippedApplicationSecrets = skippedApplicationSecrets, + skippedHandshakeSecrets = skippedHandshakeSecrets, ) } + + private fun writeSkippedSecrets( + writer: TlsWriter, + secrets: Map, ByteArray>, + ) { + writer.putUint32(secrets.size.toLong()) + for ((key, secret) in secrets) { + writer.putUint32(key.first.toLong()) + writer.putUint32(key.second.toLong()) + writer.putOpaqueVarInt(secret) + } + } + + private fun readSkippedSecrets(reader: TlsReader): Map, ByteArray> { + val count = reader.readUint32().toInt() + return buildMap { + repeat(count) { + val leafIndex = reader.readUint32().toInt() + val generation = reader.readUint32().toInt() + put(Pair(leafIndex, generation), reader.readOpaqueVarInt()) + } + } + } } } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTree.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTree.kt index 6803fb2bf3..b8366f655f 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTree.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTree.kt @@ -57,13 +57,18 @@ class SecretTree( private val consumedHandshakeGenerations = mutableMapOf>() /** - * Cache of key/nonce pairs for skipped APPLICATION generations. - * Key: (leafIndex, generation) -> derived KeyNonceGeneration. + * Ratchet secrets of skipped APPLICATION generations, for messages that + * arrive out of order. Key: (leafIndex, generation) -> that generation's + * secret; the key and nonce are derived when it is used. + * + * The secret rather than the derived pair so the cache can be persisted + * (see [exportSkippedApplicationSecrets]) in the same form other MLS + * implementations keep it (ts-mls `unusedGenerations`). */ - private val skippedKeys = mutableMapOf, KeyNonceGeneration>() + private val skippedKeys = mutableMapOf, ByteArray>() /** Same cache for the HANDSHAKE ratchet. */ - private val handshakeSkippedKeys = mutableMapOf, KeyNonceGeneration>() + private val handshakeSkippedKeys = mutableMapOf, ByteArray>() private companion object { /** Maximum number of skipped key entries to retain (prevents unbounded memory growth). */ @@ -150,15 +155,15 @@ class SecretTree( generation: Int, ): KeyNonceGeneration { // Check skipped keys cache first (out-of-order message for a previously skipped generation) - val cachedKey = skippedKeys.remove(Pair(leafIndex, generation)) - if (cachedKey != null) { + val cachedSecret = skippedKeys.remove(Pair(leafIndex, generation)) + if (cachedSecret != null) { // Still mark as consumed for replay detection val senderConsumed = consumedGenerations.getOrPut(leafIndex) { mutableSetOf() } if (generation in senderConsumed) { throw StaleGenerationException(leafIndex, generation, null, "Replay detected: generation $generation from sender $leafIndex already consumed") } senderConsumed.add(generation) - return cachedKey + return deriveKeyNonce(cachedSecret, generation) } val state = getOrInitSender(leafIndex) @@ -198,11 +203,10 @@ class SecretTree( var secret = state.applicationSecret var gen = state.applicationGeneration while (gen < generation) { - // Save the intermediate generation's key/nonce for later out-of-order retrieval - val intermediateKng = deriveKeyNonce(secret, gen) + // Save the intermediate generation's secret for later out-of-order retrieval val cacheKey = Pair(leafIndex, gen) if (skippedKeys.size < MAX_SKIPPED_KEYS) { - skippedKeys[cacheKey] = intermediateKng + skippedKeys[cacheKey] = secret } secret = MlsCryptoProvider.expandWithLabel(secret, "secret", generationContext(gen), MlsCryptoProvider.HASH_OUTPUT_LENGTH) gen++ @@ -235,8 +239,8 @@ class SecretTree( leafIndex: Int, generation: Int, ): KeyNonceGeneration { - val cachedKey = handshakeSkippedKeys.remove(Pair(leafIndex, generation)) - if (cachedKey != null) { + val cachedSecret = handshakeSkippedKeys.remove(Pair(leafIndex, generation)) + if (cachedSecret != null) { val senderConsumed = consumedHandshakeGenerations.getOrPut(leafIndex) { mutableSetOf() } if (generation in senderConsumed) { throw StaleGenerationException( @@ -247,7 +251,7 @@ class SecretTree( ) } senderConsumed.add(generation) - return cachedKey + return deriveKeyNonce(cachedSecret, generation) } val state = getOrInitSender(leafIndex) @@ -285,10 +289,9 @@ class SecretTree( var secret = state.handshakeSecret var gen = state.handshakeGeneration while (gen < generation) { - val intermediateKng = deriveKeyNonce(secret, gen) val cacheKey = Pair(leafIndex, gen) if (handshakeSkippedKeys.size < MAX_SKIPPED_KEYS) { - handshakeSkippedKeys[cacheKey] = intermediateKng + handshakeSkippedKeys[cacheKey] = secret } secret = MlsCryptoProvider.expandWithLabel(secret, "secret", generationContext(gen), MlsCryptoProvider.HASH_OUTPUT_LENGTH) gen++ @@ -435,6 +438,27 @@ class SecretTree( fun importSenderStates(states: Map) { senderState.putAll(states) } + + /** + * Secrets of skipped APPLICATION generations, keyed (leafIndex, generation). + * + * Persisted next to [exportSenderStates]: the ratchet has already moved + * past these generations, so a message that was skipped before a restart + * can only be opened after it if its secret survives. + */ + fun exportSkippedApplicationSecrets(): Map, ByteArray> = skippedKeys.toMap() + + /** Same for the HANDSHAKE ratchet. */ + fun exportSkippedHandshakeSecrets(): Map, ByteArray> = handshakeSkippedKeys.toMap() + + /** Restores skipped-generation secrets, up to the usual cache bound per ratchet. */ + fun importSkippedSecrets( + application: Map, ByteArray>, + handshake: Map, ByteArray>, + ) { + for ((k, v) in application) if (skippedKeys.size < MAX_SKIPPED_KEYS) skippedKeys[k] = v + for ((k, v) in handshake) if (handshakeSkippedKeys.size < MAX_SKIPPED_KEYS) handshakeSkippedKeys[k] = v + } } data class SenderRatchetState( diff --git a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeSkippedSecretsTest.kt b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeSkippedSecretsTest.kt new file mode 100644 index 0000000000..2e7dc21538 --- /dev/null +++ b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeSkippedSecretsTest.kt @@ -0,0 +1,58 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.quartz.mls.schedule + +import kotlin.test.Test +import kotlin.test.assertContentEquals +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +class SecretTreeSkippedSecretsTest { + private val encryptionSecret = ByteArray(32) { 3 } + + @Test + fun skippedSecretsRoundTripIntoAFreshTree() { + val reference = SecretTree(encryptionSecret, leafCount = 2) + val tree = SecretTree(encryptionSecret, leafCount = 2) + tree.applicationKeyNonceForGeneration(leafIndex = 1, generation = 3) + tree.handshakeKeyNonceForGeneration(leafIndex = 0, generation = 1) + + assertEquals(setOf(1 to 0, 1 to 1, 1 to 2), tree.exportSkippedApplicationSecrets().keys) + assertEquals(setOf(0 to 0), tree.exportSkippedHandshakeSecrets().keys) + + val restored = SecretTree(encryptionSecret, leafCount = 2) + restored.importSenderStates(tree.exportSenderStates()) + restored.importSkippedSecrets(tree.exportSkippedApplicationSecrets(), tree.exportSkippedHandshakeSecrets()) + + for (generation in 0..2) { + assertContentEquals( + reference.applicationKeyNonceForGeneration(1, generation).nonce, + restored.applicationKeyNonceForGeneration(1, generation).nonce, + ) + } + assertContentEquals( + reference.handshakeKeyNonceForGeneration(0, 0).key, + restored.handshakeKeyNonceForGeneration(0, 0).key, + ) + assertTrue(restored.exportSkippedApplicationSecrets().isEmpty()) + assertTrue(restored.exportSkippedHandshakeSecrets().isEmpty()) + } +} diff --git a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateSkippedGenerationsTest.kt b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateSkippedGenerationsTest.kt new file mode 100644 index 0000000000..56e7eeab68 --- /dev/null +++ b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateSkippedGenerationsTest.kt @@ -0,0 +1,93 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.quartz.mls + +import com.vitorpamplona.quartz.mls.group.MlsGroup +import com.vitorpamplona.quartz.mls.group.MlsGroupState +import kotlin.test.Test +import kotlin.test.assertContentEquals +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +/** + * A message that arrives after a later one from the same sender is opened + * with the secret its generation left behind. Those secrets are part of the + * saved state, so a restart between the two messages does not lose the + * earlier one. + */ +class MlsGroupStateSkippedGenerationsTest { + private fun twoMembers(): Pair { + val alice = MlsGroup.create("alice".encodeToByteArray()) + val bobBundle = + MlsGroup + .create("bob".encodeToByteArray()) + .createKeyPackage("bob".encodeToByteArray(), ByteArray(0)) + val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle) + return alice to bob + } + + private fun MlsGroup.saveAndRestore() = MlsGroup.restore(MlsGroupState.decodeTls(saveState().encodeTls())) + + @Test + fun skippedMessageOpensAfterRestore() { + val (alice, bob) = twoMembers() + val ct0 = alice.encrypt("msg0".encodeToByteArray()) + val ct1 = alice.encrypt("msg1".encodeToByteArray()) + val ct2 = alice.encrypt("msg2".encodeToByteArray()) + + // generation 2 first: bob's ratchet moves past 0 and 1 + assertContentEquals("msg2".encodeToByteArray(), bob.decrypt(ct2).content) + assertEquals(2, bob.saveState().skippedApplicationSecrets.size) + + val bobRestored = bob.saveAndRestore() + assertContentEquals("msg0".encodeToByteArray(), bobRestored.decrypt(ct0).content) + assertContentEquals("msg1".encodeToByteArray(), bobRestored.decrypt(ct1).content) + assertTrue(bobRestored.saveState().skippedApplicationSecrets.isEmpty()) + } + + @Test + fun anOpenedSkippedGenerationIsNotPersisted() { + val (alice, bob) = twoMembers() + val ct0 = alice.encrypt("msg0".encodeToByteArray()) + val ct1 = alice.encrypt("msg1".encodeToByteArray()) + + bob.decrypt(ct1) + bob.decrypt(ct0) + + val bobRestored = bob.saveAndRestore() + assertTrue(bobRestored.saveState().skippedApplicationSecrets.isEmpty()) + assertTrue(runCatching { bobRestored.decrypt(ct0) }.isFailure) + } + + @Test + fun version4StateStillDecodes() { + val (_, bob) = twoMembers() + val v5 = bob.saveState().encodeTls() + // with nothing skipped, v5 is v4 plus two empty counts (uint32 each) + val v4 = v5.copyOfRange(0, v5.size - 8) + v4[0] = 0 + v4[1] = 4 + + val restored = MlsGroup.restore(MlsGroupState.decodeTls(v4)) + assertEquals(bob.epoch, restored.epoch) + assertTrue(restored.saveState().skippedApplicationSecrets.isEmpty()) + } +} diff --git a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateTest.kt b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateTest.kt index d9b2f0e138..8717617f1a 100644 --- a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateTest.kt +++ b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateTest.kt @@ -173,14 +173,14 @@ class MlsGroupStateTest { val state = group.saveState() val bytes = state.encodeTls() - // First two bytes are the state version (uint16). v4 appends the - // staged-proposal pool; older blobs still decode, so the version only - // ever moves forward when the layout gains a field. + // First two bytes are the state version (uint16). v5 appends the + // skipped-generation secrets; older blobs still decode, so the version + // only ever moves forward when the layout gains a field. val reader = com.vitorpamplona.quartz.mls.codec .TlsReader(bytes) val version = reader.readUint16() - assertEquals(4, version) + assertEquals(5, version) } @Test