feat(mls): derive secret-tree leaves from the unexpanded node secrets

SecretTree derived every leaf from the root encryption_secret. ts-mls keeps
the tree as its unexpanded node secrets instead (intermediateNodes): each
node on the way down to a sender is replaced by its other child, so a saved
ts-mls state usually has no root left and Quartz could not load it.

SecretTree now keeps that map, starting as {root: encryption_secret}, and
derives a leaf from its nearest known ancestor. With the root present the
keys are the same as before. exportNodeSecrets/importNodeSecrets expose it,
and MlsGroupState v7 persists it for the current and retained epochs; older
blobs derive from the root as before.
This commit is contained in:
jeremyd
2026-09-29 13:09:41 -07:00
parent bee4130dd9
commit 2171b8d11d
6 changed files with 286 additions and 22 deletions
@@ -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,
@@ -100,6 +100,12 @@ data class MlsGroupState(
* after a restart. See [MlsGroup.decryptFormerEpoch].
*/
val retainedEpochs: List<RetainedEpochReceiverData> = 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<Int, ByteArray> = 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<Int, ByteArray>()
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<Int, ByteArray>,
) {
writer.putUint32(secrets.size.toLong())
for ((nodeIndex, secret) in secrets) {
writer.putUint32(nodeIndex.toLong())
writer.putOpaqueVarInt(secret)
}
}
private fun readNodeSecrets(reader: TlsReader): Map<Int, ByteArray> {
val count = reader.readUint32().toInt()
return buildMap {
repeat(count) {
val nodeIndex = reader.readUint32().toInt()
put(nodeIndex, reader.readOpaqueVarInt())
}
}
}
internal fun readSkippedSecrets(reader: TlsReader): Map<Pair<Int, Int>, ByteArray> {
val count = reader.readUint32().toInt()
return buildMap {
@@ -391,6 +435,8 @@ data class RetainedEpochReceiverData(
val treeBytes: ByteArray,
val senderRatchetStates: Map<Int, SenderRatchetState>,
val skippedApplicationSecrets: Map<Pair<Int, Int>, ByteArray>,
/** The epoch's unexpanded secret-tree node secrets; empty means derive from [encryptionSecret]. */
val nodeSecrets: Map<Int, ByteArray> = emptyMap(),
) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
@@ -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<Pair<Int, Int>, 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<Int, ByteArray>().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<Pair<Int, Int>, ByteArray> = handshakeSkippedKeys.toMap()
/** Tree-node secrets not expanded yet, by node index (ts-mls `intermediateNodes`). */
fun exportNodeSecrets(): Map<Int, ByteArray> = nodeSecrets.toMap()
/** Replaces the unexpanded node secrets, e.g. with a state saved by another implementation. */
fun importNodeSecrets(secrets: Map<Int, ByteArray>) {
nodeSecrets.clear()
nodeSecrets.putAll(secrets)
}
/** Restores skipped-generation secrets, up to the usual cache bound per ratchet. */
fun importSkippedSecrets(
application: Map<Pair<Int, Int>, ByteArray>,
@@ -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<IllegalArgumentException> { key(imported, 2) }
}
}
@@ -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<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
}
@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)
}
}
@@ -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