Merge pull request #4271 from jeremyd/feat/quartz-mls-persist-skipped-generations

feat(mls): keep skipped-generation secrets across a restore
This commit is contained in:
Vitor Pamplona
2026-09-29 17:39:31 -04:00
committed by GitHub
6 changed files with 245 additions and 20 deletions
@@ -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,
@@ -83,6 +83,17 @@ data class MlsGroupState(
* the proposal that would evict them.
*/
val pendingProposals: List<PendingProposal> = 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<Pair<Int, Int>, ByteArray> = emptyMap(),
/** Same for the HANDSHAKE ratchet (STATE_VERSION 5+). */
val skippedHandshakeSecrets: Map<Pair<Int, Int>, 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<Pair<Int, Int>, 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<Pair<Int, Int>, 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())
}
}
}
}
}
@@ -57,13 +57,18 @@ class SecretTree(
private val consumedHandshakeGenerations = mutableMapOf<Int, MutableSet<Int>>()
/**
* 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<Pair<Int, Int>, KeyNonceGeneration>()
private val skippedKeys = mutableMapOf<Pair<Int, Int>, ByteArray>()
/** Same cache for the HANDSHAKE ratchet. */
private val handshakeSkippedKeys = mutableMapOf<Pair<Int, Int>, KeyNonceGeneration>()
private val handshakeSkippedKeys = mutableMapOf<Pair<Int, Int>, 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<Int, SenderRatchetState>) {
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<Pair<Int, Int>, ByteArray> = skippedKeys.toMap()
/** Same for the HANDSHAKE ratchet. */
fun exportSkippedHandshakeSecrets(): Map<Pair<Int, Int>, ByteArray> = handshakeSkippedKeys.toMap()
/** Restores skipped-generation secrets, up to the usual cache bound per ratchet. */
fun importSkippedSecrets(
application: Map<Pair<Int, Int>, ByteArray>,
handshake: Map<Pair<Int, Int>, 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(
@@ -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())
}
}
@@ -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<MlsGroup, MlsGroup> {
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())
}
}
@@ -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