mirror of
https://github.com/vitorpamplona/amethyst.git
synced 2026-10-05 19:28:25 +00:00
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:
@@ -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(
|
||||
|
||||
+58
@@ -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())
|
||||
}
|
||||
}
|
||||
+93
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user