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 591a9d8131..083cf18c08 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 @@ -295,6 +295,7 @@ class MlsGroup private constructor( skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(), skippedHandshakeSecrets = secretTree.exportSkippedHandshakeSecrets(), retainedEpochs = retainedEpochs.toList(), + nodeSecrets = secretTree.exportNodeSecrets(), ) } @@ -415,6 +416,7 @@ class MlsGroup private constructor( treeBytes = w.toByteArray(), senderRatchetStates = secretTree.exportSenderStates(), skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(), + nodeSecrets = secretTree.exportNodeSecrets(), ) } @@ -1500,6 +1502,7 @@ class MlsGroup private constructor( val formerSecrets = SecretTree(retained.encryptionSecret, formerTree.leafCount) formerSecrets.importSenderStates(retained.senderRatchetStates) formerSecrets.importSkippedSecrets(retained.skippedApplicationSecrets, emptyMap()) + if (retained.nodeSecrets.isNotEmpty()) formerSecrets.importNodeSecrets(retained.nodeSecrets) val opened = openPrivateMessage(privMsg, retained.senderDataSecret, formerTree, formerSecrets, allowOwnSentKeys = false) val decrypted = applicationFromOpened(privMsg, opened, formerTree, retained.groupContext) @@ -1507,6 +1510,7 @@ class MlsGroup private constructor( retained.copy( senderRatchetStates = formerSecrets.exportSenderStates(), skippedApplicationSecrets = formerSecrets.exportSkippedApplicationSecrets(), + nodeSecrets = formerSecrets.exportNodeSecrets(), ) return decrypted } @@ -4158,6 +4162,7 @@ class MlsGroup private constructor( val secretTree = SecretTree(state.encryptionSecret, tree.leafCount) secretTree.importSenderStates(state.senderRatchetStates) secretTree.importSkippedSecrets(state.skippedApplicationSecrets, state.skippedHandshakeSecrets) + if (state.nodeSecrets.isNotEmpty()) secretTree.importNodeSecrets(state.nodeSecrets) 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 92a1b9e7ba..68fa450dd7 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 @@ -100,6 +100,12 @@ data class MlsGroupState( * after a restart. See [MlsGroup.decryptFormerEpoch]. */ val retainedEpochs: List = emptyList(), + /** + * The secret tree's unexpanded node secrets (STATE_VERSION 7+). Empty + * means "derive from [encryptionSecret]", which is what every older blob + * does. A state imported from ts-mls has these instead of a root. + */ + val nodeSecrets: Map = emptyMap(), ) { fun encodeTls(): ByteArray { val writer = TlsWriter() @@ -181,6 +187,11 @@ data class MlsGroupState( writer.putUint32(retainedEpochs.size.toLong()) for (retained in retainedEpochs) retained.encodeTls(writer) + // Secret-tree node secrets (STATE_VERSION 7+), the current epoch's, + // then each retained epoch's in the same order as above. + writeNodeSecrets(writer, nodeSecrets) + for (retained in retainedEpochs) writeNodeSecrets(writer, retained.nodeSecrets) + return writer.toByteArray() } @@ -212,8 +223,11 @@ data class MlsGroupState( * decode with none, as before. * v6: appends [retainedEpochs]. Older blobs decode with none, so the * first commit after the upgrade starts the window. + * v7: appends the secret tree's [nodeSecrets], current and retained + * epochs, so a state whose root secret is gone (imported from + * ts-mls) restores. Older blobs derive from the root, as before. */ - private const val STATE_VERSION = 6 + private const val STATE_VERSION = 7 fun decodeTls(data: ByteArray): MlsGroupState { val reader = TlsReader(data) @@ -327,6 +341,14 @@ data class MlsGroupState( emptyList() } + // v7+: secret-tree node secrets. Absent for older blobs. + var nodeSecrets = emptyMap() + var retainedWithNodes = retainedEpochs + if (version >= 7 && reader.hasRemaining) { + nodeSecrets = readNodeSecrets(reader) + retainedWithNodes = retainedEpochs.map { it.copy(nodeSecrets = readNodeSecrets(reader)) } + } + return MlsGroupState( groupContext = groupContext, treeBytes = treeBytes, @@ -342,7 +364,8 @@ data class MlsGroupState( pendingProposals = pendingProposals, skippedApplicationSecrets = skippedApplicationSecrets, skippedHandshakeSecrets = skippedHandshakeSecrets, - retainedEpochs = retainedEpochs, + retainedEpochs = retainedWithNodes, + nodeSecrets = nodeSecrets, ) } @@ -358,6 +381,27 @@ data class MlsGroupState( } } + private fun writeNodeSecrets( + writer: TlsWriter, + secrets: Map, + ) { + writer.putUint32(secrets.size.toLong()) + for ((nodeIndex, secret) in secrets) { + writer.putUint32(nodeIndex.toLong()) + writer.putOpaqueVarInt(secret) + } + } + + private fun readNodeSecrets(reader: TlsReader): Map { + val count = reader.readUint32().toInt() + return buildMap { + repeat(count) { + val nodeIndex = reader.readUint32().toInt() + put(nodeIndex, reader.readOpaqueVarInt()) + } + } + } + internal fun readSkippedSecrets(reader: TlsReader): Map, ByteArray> { val count = reader.readUint32().toInt() return buildMap { @@ -391,6 +435,8 @@ data class RetainedEpochReceiverData( val treeBytes: ByteArray, val senderRatchetStates: Map, val skippedApplicationSecrets: Map, ByteArray>, + /** The epoch's unexpanded secret-tree node secrets; empty means derive from [encryptionSecret]. */ + val nodeSecrets: Map = emptyMap(), ) { override fun equals(other: Any?): Boolean { if (this === other) return true 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 b8366f655f..75f9b3dd47 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 @@ -44,7 +44,7 @@ import com.vitorpamplona.quartz.mls.tree.BinaryTree * ``` */ class SecretTree( - private val encryptionSecret: ByteArray, + encryptionSecret: ByteArray, private val leafCount: Int, ) { /** Per-sender ratchet state: (handshake generation, handshake secret, app generation, app secret) */ @@ -70,6 +70,22 @@ class SecretTree( /** Same cache for the HANDSHAKE ratchet. */ private val handshakeSkippedKeys = mutableMapOf, ByteArray>() + /** + * Tree-node secrets not expanded yet, by node index. Starts as + * `{root: encryptionSecret}`. Deriving a leaf replaces each node on the + * way down with its other child, the way RFC 9420 §9 describes and ts-mls + * keeps it (`intermediateNodes`). + * + * A state imported from another implementation usually no longer has the + * root: once a sender has been derived, only the siblings along its path + * are left. [importNodeSecrets] takes that map and later senders are + * derived from their nearest known ancestor. + */ + private val nodeSecrets = + mutableMapOf().also { + if (encryptionSecret.isNotEmpty()) it[BinaryTree.root(leafCount)] = encryptionSecret + } + private companion object { /** Maximum number of skipped key entries to retain (prevents unbounded memory growth). */ const val MAX_SKIPPED_KEYS = 1000 @@ -366,8 +382,9 @@ class SecretTree( } /** - * Derive the leaf secret from the encryption secret by walking DOWN the - * binary tree from the root to the target leaf (RFC 9420 §9). + * Derive the leaf secret by walking DOWN the binary tree from the nearest + * ancestor in [nodeSecrets] (the root, in a fresh tree) to the target + * leaf (RFC 9420 §9). * * At each step we pick left or right based on which subtree contains the * target. In an MLS left-balanced tree the left-subtree node indices are @@ -393,23 +410,32 @@ class SecretTree( */ private fun getLeafSecret(leafIndex: Int): ByteArray { val targetNode = BinaryTree.leafToNode(leafIndex) - val rootIdx = BinaryTree.root(leafCount) - var currentSecret = encryptionSecret - var currentNode = rootIdx + nodeSecrets.remove(targetNode)?.let { return it } - while (currentNode != targetNode) { - val goLeft = targetNode < currentNode - val label = if (goLeft) "left" else "right" - currentSecret = - MlsCryptoProvider.expandWithLabel( - currentSecret, - "tree", - label.encodeToByteArray(), - MlsCryptoProvider.HASH_OUTPUT_LENGTH, - ) - currentNode = if (goLeft) BinaryTree.left(currentNode) else BinaryTree.right(currentNode) + val path = mutableListOf(BinaryTree.root(leafCount)) + while (path.last() != targetNode) { + val node = path.last() + path.add(if (targetNode < node) BinaryTree.left(node) else BinaryTree.right(node)) } + // Start from the nearest ancestor whose secret is still known. + val start = path.indexOfLast { it in nodeSecrets } + require(start >= 0) { "No secret-tree node secret left to derive leaf $leafIndex" } + + var currentSecret = nodeSecrets.remove(path[start])!! + for (k in start until path.size - 1) { + val node = path[k] + val left = MlsCryptoProvider.expandWithLabel(currentSecret, "tree", "left".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH) + val right = MlsCryptoProvider.expandWithLabel(currentSecret, "tree", "right".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH) + // keep the other child's secret for the senders under it + if (path[k + 1] == BinaryTree.left(node)) { + nodeSecrets[BinaryTree.right(node)] = right + currentSecret = left + } else { + nodeSecrets[BinaryTree.left(node)] = left + currentSecret = right + } + } return currentSecret } @@ -451,6 +477,15 @@ class SecretTree( /** Same for the HANDSHAKE ratchet. */ fun exportSkippedHandshakeSecrets(): Map, ByteArray> = handshakeSkippedKeys.toMap() + /** Tree-node secrets not expanded yet, by node index (ts-mls `intermediateNodes`). */ + fun exportNodeSecrets(): Map = nodeSecrets.toMap() + + /** Replaces the unexpanded node secrets, e.g. with a state saved by another implementation. */ + fun importNodeSecrets(secrets: Map) { + nodeSecrets.clear() + nodeSecrets.putAll(secrets) + } + /** Restores skipped-generation secrets, up to the usual cache bound per ratchet. */ fun importSkippedSecrets( application: Map, ByteArray>, diff --git a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeNodeSecretsTest.kt b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeNodeSecretsTest.kt new file mode 100644 index 0000000000..265b89642f --- /dev/null +++ b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeNodeSecretsTest.kt @@ -0,0 +1,82 @@ +/* + * 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.assertFailsWith + +class SecretTreeNodeSecretsTest { + private val encryptionSecret = ByteArray(32) { 9 } + + private fun key( + tree: SecretTree, + leafIndex: Int, + ) = tree.applicationKeyNonceForGeneration(leafIndex, 0).key + + @Test + fun derivingALeafKeepsTheSiblingsAndDropsThePath() { + // 4 leaves: leaf nodes 0, 2, 4, 6; parents 1 and 5; root 3 + val tree = SecretTree(encryptionSecret, leafCount = 4) + assertEquals(setOf(3), tree.exportNodeSecrets().keys) + + key(tree, 0) + assertEquals(setOf(2, 5), tree.exportNodeSecrets().keys) + + key(tree, 3) + assertEquals(setOf(2, 4), tree.exportNodeSecrets().keys) + } + + @Test + fun aTreeWithoutItsRootDerivesTheSameKeys() { + val reference = SecretTree(encryptionSecret, leafCount = 4) + val source = SecretTree(encryptionSecret, leafCount = 4) + key(source, 0) + + // what another implementation hands over once leaf 0 has sent: no root + val imported = SecretTree(ByteArray(0), leafCount = 4) + imported.importNodeSecrets(source.exportNodeSecrets()) + + for (leaf in 1..3) assertContentEquals(key(reference, leaf), key(imported, leaf)) + } + + @Test + fun nonPowerOfTwoTreesDeriveTheSameKeys() { + val reference = SecretTree(encryptionSecret, leafCount = 5) + val source = SecretTree(encryptionSecret, leafCount = 5) + key(source, 4) + + val imported = SecretTree(ByteArray(0), leafCount = 5) + imported.importNodeSecrets(source.exportNodeSecrets()) + for (leaf in 0..3) assertContentEquals(key(reference, leaf), key(imported, leaf)) + } + + @Test + fun aLeafWithNoKnownAncestorFails() { + val source = SecretTree(encryptionSecret, leafCount = 4) + key(source, 0) + val imported = SecretTree(ByteArray(0), leafCount = 4) + imported.importNodeSecrets(source.exportNodeSecrets().filterKeys { it != 5 }) + + assertFailsWith { key(imported, 2) } + } +} diff --git a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateNodeSecretsTest.kt b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateNodeSecretsTest.kt new file mode 100644 index 0000000000..5c4194f02b --- /dev/null +++ b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateNodeSecretsTest.kt @@ -0,0 +1,96 @@ +/* + * 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 saved state without the secret tree's root, as another implementation + * (ts-mls) keeps it, restores from the unexpanded node secrets alone. + */ +class MlsGroupStateNodeSecretsTest { + 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 + } + + @Test + fun aStateWithoutTheRootSecretRestores() { + val (alice, bob) = twoMembers() + // bob derives alice's leaf, so his tree no longer needs the root + assertContentEquals("hi".encodeToByteArray(), bob.decrypt(alice.encrypt("hi".encodeToByteArray())).content) + + val saved = bob.saveState() + assertTrue(saved.nodeSecrets.isNotEmpty()) + val rootless = saved.copy(encryptionSecret = ByteArray(0)) + val restored = MlsGroup.restore(MlsGroupState.decodeTls(rootless.encodeTls())) + + // bob's own leaf comes from the kept sibling, and alice can read it + assertContentEquals("from bob".encodeToByteArray(), alice.decrypt(restored.encrypt("from bob".encodeToByteArray())).content) + // alice's ratchet continues where it was + assertContentEquals("again".encodeToByteArray(), restored.decrypt(alice.encrypt("again".encodeToByteArray())).content) + } + + @Test + fun retainedEpochsKeepTheirNodeSecrets() { + val (alice, bob) = twoMembers() + val late = bob.encrypt("late".encodeToByteArray()) + alice.decrypt(bob.encrypt("seen".encodeToByteArray())) + alice.commit() + + val retained = alice.retainedEpochs().last() + assertTrue(retained.nodeSecrets.isNotEmpty()) + val restored = MlsGroup.restore(MlsGroupState.decodeTls(alice.saveState().encodeTls())) + assertEquals( + retained.nodeSecrets.keys, + restored + .retainedEpochs() + .last() + .nodeSecrets.keys, + ) + assertContentEquals("late".encodeToByteArray(), restored.decryptFormerEpoch(late).content) + } + + @Test + fun version6StateStillDecodes() { + val (_, bob) = twoMembers() + val state = bob.saveState() + val v7 = state.copy(nodeSecrets = emptyMap()).encodeTls() + // with no retained epochs and no node secrets, v7 is v6 plus one empty count (uint32) + assertTrue(state.retainedEpochs.isEmpty()) + val v6 = v7.copyOfRange(0, v7.size - 4) + v6[0] = 0 + v6[1] = 6 + + val restored = MlsGroup.restore(MlsGroupState.decodeTls(v6)) + assertEquals(bob.epoch, restored.epoch) + } +} 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 1013e054ab..2d869abd82 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). v6 appends the - // retained former epochs; older blobs still decode, so the version + // First two bytes are the state version (uint16). v7 appends the + // secret-tree node 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(6, version) + assertEquals(7, version) } @Test