Merge pull request #4275 from jeremyd/feat/quartz-mls-secret-tree-node-secrets

feat(mls): derive secret-tree leaves from the unexpanded node secrets
This commit is contained in:
Vitor Pamplona
2026-09-29 17:39:50 -04:00
committed by GitHub
7 changed files with 716 additions and 81 deletions
@@ -125,6 +125,8 @@ class MlsGroup private constructor(
private var interimTranscriptHash: ByteArray,
private val pskStore: MutableMap<String, ByteArray> = mutableMapOf(),
private val pendingProposals: MutableList<PendingProposal> = mutableListOf(),
/** Receiver data of the last [RETAIN_EPOCHS] epochs, oldest first. See [decryptFormerEpoch]. */
private val retainedEpochs: ArrayDeque<RetainedEpochReceiverData> = ArrayDeque(),
private val sentKeys: MutableMap<Int, com.vitorpamplona.quartz.mls.schedule.KeyNonceGeneration> = mutableMapOf(),
/** Staged keys from proposeSigningKeyRotation — only promoted on successful commit */
private var pendingSigningKey: ByteArray? = null,
@@ -292,6 +294,8 @@ class MlsGroup private constructor(
pendingProposals = pendingProposals.toList(),
skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(),
skippedHandshakeSecrets = secretTree.exportSkippedHandshakeSecrets(),
retainedEpochs = retainedEpochs.toList(),
nodeSecrets = secretTree.exportNodeSecrets(),
)
}
@@ -398,6 +402,33 @@ class MlsGroup private constructor(
pathPrivateKeys.keys.retainAll(fullPath.toSet())
}
/** The current epoch's receiver data. Taken before a commit changes anything, kept once it has applied. */
private fun captureRetainedEpoch(): RetainedEpochReceiverData {
val w = TlsWriter()
tree.encodeTls(w)
return RetainedEpochReceiverData(
epoch = epoch,
senderDataSecret = epochSecrets.senderDataSecret,
encryptionSecret = epochSecrets.encryptionSecret,
exporterSecret = epochSecrets.exporterSecret,
resumptionPsk = epochSecrets.resumptionPsk,
groupContext = groupContext,
treeBytes = w.toByteArray(),
senderRatchetStates = secretTree.exportSenderStates(),
skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(),
nodeSecrets = secretTree.exportNodeSecrets(),
)
}
private fun pushRetainedEpoch(retained: RetainedEpochReceiverData) {
retainedEpochs.removeAll { it.epoch == retained.epoch }
retainedEpochs.addLast(retained)
while (retainedEpochs.size > RETAIN_EPOCHS) retainedEpochs.removeFirst()
}
/** Receiver data of the retained former epochs, oldest first. */
fun retainedEpochs(): List<RetainedEpochReceiverData> = retainedEpochs.toList()
/**
* Extract retained epoch secrets for late-message decryption.
*
@@ -669,6 +700,7 @@ class MlsGroup private constructor(
* Returns the Commit bytes to send to the group, plus optional Welcome for new members.
*/
fun commit(): CommitResult {
val retainedForThisEpoch = captureRetainedEpoch()
val proposals = pendingProposals.toList()
// The application's gate on who may commit what. RFC 9420 has none of
@@ -1028,6 +1060,7 @@ class MlsGroup private constructor(
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, initSecret, pskSecret)
initSecret = epochSecrets.initSecret
secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
pushRetainedEpoch(retainedForThisEpoch)
// Compute confirmation_tag and interim_transcript_hash
val confirmationTag = computeConfirmationTag(epochSecrets.confirmationKey, newConfirmedTranscriptHash)
@@ -1307,25 +1340,23 @@ class MlsGroup private constructor(
null
}
/** The AEAD-opened content of a PrivateMessage: who sent it and the PrivateMessageContent bytes. */
private class OpenedPrivateMessage(
val senderLeafIndex: Int,
val plaintext: ByteArray,
)
/**
* Decrypt an application message from a PrivateMessage (RFC 9420 Section 6.3).
* @throws IllegalArgumentException if the message format is invalid
* @throws javax.crypto.AEADBadTagException if decryption fails
* Sender-data decryption, ratchet key lookup and content AEAD for a PrivateMessage against the given
* epoch material (RFC 9420 §6.3). Shared by [decrypt] (current epoch) and [decryptFormerEpoch].
*/
fun decrypt(messageBytes: ByteArray): DecryptedMessage {
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
// Verify epoch and group ID match current state (RFC 9420 Section 6.1)
require(privMsg.epoch == epoch) {
"Message epoch ${privMsg.epoch} doesn't match current epoch $epoch"
}
require(privMsg.groupId.contentEquals(groupId)) {
"Message group ID doesn't match current group"
}
private fun openPrivateMessage(
privMsg: PrivateMessage,
senderDataSecret: ByteArray,
tree: RatchetTree,
secretTree: SecretTree,
allowOwnSentKeys: Boolean,
): OpenedPrivateMessage {
// Derive sender data key/nonce using ciphertext sample (RFC 9420 §6.3.1)
// RFC 9420 §6.3.2: ciphertext_sample is the first KDF.Nh bytes
// (32 for HKDF-SHA256), not AEAD.Nk (16). Using AEAD.Nk here made
@@ -1334,14 +1365,14 @@ class MlsGroup private constructor(
privMsg.ciphertext.copyOfRange(0, minOf(privMsg.ciphertext.size, MlsCryptoProvider.HASH_OUTPUT_LENGTH))
val senderDataKey =
MlsCryptoProvider.expandWithLabel(
epochSecrets.senderDataSecret,
senderDataSecret,
"key",
ciphertextSample,
MlsCryptoProvider.AEAD_KEY_LENGTH,
)
val senderDataNonce =
MlsCryptoProvider.expandWithLabel(
epochSecrets.senderDataSecret,
senderDataSecret,
"nonce",
ciphertextSample,
MlsCryptoProvider.AEAD_NONCE_LENGTH,
@@ -1371,7 +1402,7 @@ class MlsGroup private constructor(
// outgoing, so every B→A commit quartz receives lands here with
// content_type == COMMIT.
val kng =
if (senderLeafIndex == myLeafIndex && sentKeys.containsKey(generation)) {
if (allowOwnSentKeys && senderLeafIndex == myLeafIndex && sentKeys.containsKey(generation)) {
sentKeys.remove(generation)!!
} else {
when (privMsg.contentType) {
@@ -1395,49 +1426,126 @@ class MlsGroup private constructor(
val contentAad = buildPrivateContentAAD(privMsg.groupId, privMsg.epoch, privMsg.contentType, privMsg.authenticatedData)
val pmcPlaintext = MlsCryptoProvider.aeadDecrypt(kng.key, guardedNonce, contentAad, privMsg.ciphertext)
return OpenedPrivateMessage(senderLeafIndex, pmcPlaintext)
}
/** Parses and signature-checks an APPLICATION PrivateMessageContent against the given tree/context. */
private fun applicationFromOpened(
privMsg: PrivateMessage,
opened: OpenedPrivateMessage,
tree: RatchetTree,
groupContext: GroupContext,
): DecryptedMessage {
val pmcReader = TlsReader(opened.plaintext)
val senderLeafIndex = opened.senderLeafIndex
val applicationData = pmcReader.readOpaqueVarInt()
val signature = pmcReader.readOpaqueVarInt()
while (pmcReader.hasRemaining) {
require(pmcReader.readBytes(1)[0] == 0.toByte()) {
"PrivateMessageContent padding must be zero"
}
}
val senderLeaf =
requireNotNull(tree.getLeaf(senderLeafIndex)) {
"Sender leaf is blank at index $senderLeafIndex"
}
require(
MlsCryptoProvider.verifyWithLabel(
senderLeaf.signatureKey,
"FramedContentTBS",
buildApplicationFramedContentTbs(
groupId = privMsg.groupId,
epoch = privMsg.epoch,
senderLeafIndex = senderLeafIndex,
authenticatedData = privMsg.authenticatedData,
applicationData = applicationData,
groupContext = groupContext,
),
signature,
),
) { "FramedContentTBS signature verification failed" }
return DecryptedMessage(
senderLeafIndex = senderLeafIndex,
contentType = privMsg.contentType,
content = applicationData,
epoch = privMsg.epoch,
authenticatedData = privMsg.authenticatedData,
)
}
/**
* Opens an APPLICATION message sealed in a retained former epoch: one from
* a member that had not yet seen the latest commit (RFC 9420 §15.2 asks
* receivers to keep recent epochs' keys for exactly this).
*
* Same checks as [decrypt] against that epoch's tree and GroupContext,
* signature included, and the epoch's ratchet is written back so a
* generation opens only once. Handshake messages from a former epoch are
* refused, as is an epoch no longer retained.
*/
fun decryptFormerEpoch(messageBytes: ByteArray): DecryptedMessage {
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
require(privMsg.epoch < epoch) { "Message epoch ${privMsg.epoch} is not a former epoch (current $epoch)" }
require(privMsg.contentType == ContentType.APPLICATION) { "Only application messages can be read from a former epoch" }
require(privMsg.groupId.contentEquals(groupId)) { "Message group ID doesn't match current group" }
val index = retainedEpochs.indexOfFirst { it.epoch == privMsg.epoch }
require(index >= 0) {
"Message epoch ${privMsg.epoch} is no longer retained (oldest kept: ${retainedEpochs.firstOrNull()?.epoch ?: "none"})"
}
val retained = retainedEpochs[index]
val formerTree = RatchetTree.decodeTls(TlsReader(retained.treeBytes))
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)
retainedEpochs[index] =
retained.copy(
senderRatchetStates = formerSecrets.exportSenderStates(),
skippedApplicationSecrets = formerSecrets.exportSkippedApplicationSecrets(),
nodeSecrets = formerSecrets.exportNodeSecrets(),
)
return decrypted
}
/** Exporter secrets of the retained former epochs, by epoch, for keys an application derives per epoch. */
fun formerExporterSecrets(): Map<Long, ByteArray> = retainedEpochs.associate { it.epoch to it.exporterSecret }
/**
* Decrypt an application message from a PrivateMessage (RFC 9420 Section 6.3).
* @throws IllegalArgumentException if the message format is invalid
* @throws javax.crypto.AEADBadTagException if decryption fails
*/
fun decrypt(messageBytes: ByteArray): DecryptedMessage {
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
// Verify epoch and group ID match current state (RFC 9420 Section 6.1)
require(privMsg.epoch == epoch) {
"Message epoch ${privMsg.epoch} doesn't match current epoch $epoch"
}
require(privMsg.groupId.contentEquals(groupId)) {
"Message group ID doesn't match current group"
}
val opened = openPrivateMessage(privMsg, epochSecrets.senderDataSecret, tree, secretTree, allowOwnSentKeys = true)
val senderLeafIndex = opened.senderLeafIndex
// Parse PrivateMessageContent (RFC 9420 §6.3.1). The layout depends on
// content_type — application payloads carry `opaque application_data<V>`
// whereas commit / proposal payloads carry the struct directly (no
// outer length prefix).
val pmcReader = TlsReader(pmcPlaintext)
val pmcReader = TlsReader(opened.plaintext)
when (privMsg.contentType) {
ContentType.APPLICATION -> {
val applicationData = pmcReader.readOpaqueVarInt()
val signature = pmcReader.readOpaqueVarInt()
while (pmcReader.hasRemaining) {
require(pmcReader.readBytes(1)[0] == 0.toByte()) {
"PrivateMessageContent padding must be zero"
}
}
val senderLeaf =
requireNotNull(tree.getLeaf(senderLeafIndex)) {
"Sender leaf is blank at index $senderLeafIndex"
}
require(
MlsCryptoProvider.verifyWithLabel(
senderLeaf.signatureKey,
"FramedContentTBS",
buildApplicationFramedContentTbs(
groupId = privMsg.groupId,
epoch = privMsg.epoch,
senderLeafIndex = senderLeafIndex,
authenticatedData = privMsg.authenticatedData,
applicationData = applicationData,
groupContext = groupContext,
),
signature,
),
) { "FramedContentTBS signature verification failed" }
return DecryptedMessage(
senderLeafIndex = senderLeafIndex,
contentType = privMsg.contentType,
content = applicationData,
epoch = privMsg.epoch,
authenticatedData = privMsg.authenticatedData,
)
}
ContentType.APPLICATION -> return applicationFromOpened(privMsg, opened, tree, groupContext)
ContentType.COMMIT -> {
// PrivateMessageContent for a Commit: the Commit struct
@@ -1691,6 +1799,7 @@ class MlsGroup private constructor(
val interimSnapshot = interimTranscriptHash
val pendingSnapshot = pendingProposals.toList()
val sentKeysSnapshot = sentKeys.toMap()
val retainedSnapshot = retainedEpochs.toList()
try {
processCommitInner(commitBytes, senderLeafIndex, confirmationTag, signature, wireFormat)
@@ -1705,6 +1814,8 @@ class MlsGroup private constructor(
pendingProposals.addAll(pendingSnapshot)
sentKeys.clear()
sentKeys.putAll(sentKeysSnapshot)
retainedEpochs.clear()
retainedEpochs.addAll(retainedSnapshot)
throw t
}
}
@@ -1716,6 +1827,7 @@ class MlsGroup private constructor(
signature: ByteArray,
wireFormat: WireFormat,
) {
val retainedForThisEpoch = captureRetainedEpoch()
val commit = Commit.decodeTls(TlsReader(commitBytes))
// External commits (containing ExternalInit) have a sender that is not
@@ -2109,6 +2221,7 @@ class MlsGroup private constructor(
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, effectiveInitSecret, pskSecret)
initSecret = epochSecrets.initSecret
secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
pushRetainedEpoch(retainedForThisEpoch)
// Verify confirmation tag (RFC 9420 Section 6.1). Every commit on
// the wire MUST carry a confirmation_tag that matches what the
@@ -3224,6 +3337,13 @@ class MlsGroup private constructor(
/** MLS extensions draft `app_data_update` proposal type. */
const val APP_DATA_UPDATE_PROPOSAL_TYPE = 0x0008
/**
* How many former epochs [decryptFormerEpoch] can still read. Each one
* keeps that epoch's decryption secrets, so this is a forward-secrecy
* trade (RFC 9420 §15.2); 4 matches ts-mls's `retainKeysForEpochs`.
*/
const val RETAIN_EPOCHS = 4
/** How far back a fresh KeyPackage LeafNode's `not_before` is set. */
private const val LIFETIME_SKEW_SECONDS = 3_600L
@@ -4042,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,
@@ -4055,6 +4176,7 @@ class MlsGroup private constructor(
interimTranscriptHash = state.interimTranscriptHash,
pathPrivateKeys = state.pathPrivateKeys.toMutableMap(),
pendingProposals = state.pendingProposals.toMutableList(),
retainedEpochs = ArrayDeque(state.retainedEpochs),
policy = policy,
)
}
@@ -94,6 +94,18 @@ data class MlsGroupState(
val skippedApplicationSecrets: Map<Pair<Int, Int>, ByteArray> = emptyMap(),
/** Same for the HANDSHAKE ratchet (STATE_VERSION 5+). */
val skippedHandshakeSecrets: Map<Pair<Int, Int>, ByteArray> = emptyMap(),
/**
* Receiver data of the last [MlsGroup.RETAIN_EPOCHS] epochs, oldest first
* (STATE_VERSION 6+), so a late message from a former epoch still opens
* 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()
@@ -171,6 +183,15 @@ data class MlsGroupState(
writeSkippedSecrets(writer, skippedApplicationSecrets)
writeSkippedSecrets(writer, skippedHandshakeSecrets)
// Retained former epochs (STATE_VERSION 6+).
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()
}
@@ -200,8 +221,13 @@ data class MlsGroupState(
* v5: appends [skippedApplicationSecrets] and [skippedHandshakeSecrets]
* so an out-of-order message survives a restart. Older blobs
* 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 = 5
private const val STATE_VERSION = 7
fun decodeTls(data: ByteArray): MlsGroupState {
val reader = TlsReader(data)
@@ -306,6 +332,23 @@ data class MlsGroupState(
val skippedApplicationSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap()
val skippedHandshakeSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap()
// v6+: retained former epochs. Absent for older blobs.
val retainedEpochs =
if (version >= 6 && reader.hasRemaining) {
val count = reader.readUint32().toInt()
List(count) { RetainedEpochReceiverData.decodeTls(reader) }
} else {
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,
@@ -321,10 +364,12 @@ data class MlsGroupState(
pendingProposals = pendingProposals,
skippedApplicationSecrets = skippedApplicationSecrets,
skippedHandshakeSecrets = skippedHandshakeSecrets,
retainedEpochs = retainedWithNodes,
nodeSecrets = nodeSecrets,
)
}
private fun writeSkippedSecrets(
internal fun writeSkippedSecrets(
writer: TlsWriter,
secrets: Map<Pair<Int, Int>, ByteArray>,
) {
@@ -336,7 +381,28 @@ data class MlsGroupState(
}
}
private fun readSkippedSecrets(reader: TlsReader): Map<Pair<Int, Int>, ByteArray> {
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 {
repeat(count) {
@@ -349,6 +415,92 @@ data class MlsGroupState(
}
}
/**
* What [MlsGroup.decryptFormerEpoch] needs to open a late APPLICATION message
* from one former epoch: its sender-data and encryption secrets, the sender
* ratchets and skipped generations as the epoch left them, and the tree and
* GroupContext the sender's signature is checked against.
*
* The exporter secret and resumption PSK ride along so keys an application
* derived in that epoch, and a resumption PSK for it (RFC 9420 §8.6), stay
* available for as long as the epoch is retained.
*/
data class RetainedEpochReceiverData(
val epoch: Long,
val senderDataSecret: ByteArray,
val encryptionSecret: ByteArray,
val exporterSecret: ByteArray,
val resumptionPsk: ByteArray,
val groupContext: GroupContext,
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
if (other !is RetainedEpochReceiverData) return false
return epoch == other.epoch
}
override fun hashCode(): Int = epoch.hashCode()
fun encodeTls(writer: TlsWriter) {
writer.putUint64(epoch)
writer.putOpaqueVarInt(senderDataSecret)
writer.putOpaqueVarInt(encryptionSecret)
writer.putOpaqueVarInt(exporterSecret)
writer.putOpaqueVarInt(resumptionPsk)
writer.putOpaqueVarInt(groupContext.toTlsBytes())
writer.putOpaqueVarInt(treeBytes)
writer.putUint32(senderRatchetStates.size.toLong())
for ((leafIndex, ratchet) in senderRatchetStates) {
writer.putUint32(leafIndex.toLong())
writer.putOpaqueVarInt(ratchet.handshakeSecret)
writer.putUint32(ratchet.handshakeGeneration.toLong())
writer.putOpaqueVarInt(ratchet.applicationSecret)
writer.putUint32(ratchet.applicationGeneration.toLong())
}
MlsGroupState.writeSkippedSecrets(writer, skippedApplicationSecrets)
}
companion object {
fun decodeTls(reader: TlsReader): RetainedEpochReceiverData {
val epoch = reader.readUint64()
val senderDataSecret = reader.readOpaqueVarInt()
val encryptionSecret = reader.readOpaqueVarInt()
val exporterSecret = reader.readOpaqueVarInt()
val resumptionPsk = reader.readOpaqueVarInt()
val groupContext = GroupContext.decodeTls(TlsReader(reader.readOpaqueVarInt()))
val treeBytes = reader.readOpaqueVarInt()
val count = reader.readUint32().toInt()
val senderRatchetStates =
buildMap {
repeat(count) {
val leafIndex = reader.readUint32().toInt()
val handshakeSecret = reader.readOpaqueVarInt()
val handshakeGeneration = reader.readUint32().toInt()
val applicationSecret = reader.readOpaqueVarInt()
val applicationGeneration = reader.readUint32().toInt()
put(leafIndex, SenderRatchetState(handshakeSecret, handshakeGeneration, applicationSecret, applicationGeneration))
}
}
return RetainedEpochReceiverData(
epoch = epoch,
senderDataSecret = senderDataSecret,
encryptionSecret = encryptionSecret,
exporterSecret = exporterSecret,
resumptionPsk = resumptionPsk,
groupContext = groupContext,
treeBytes = treeBytes,
senderRatchetStates = senderRatchetStates,
skippedApplicationSecrets = MlsGroupState.readSkippedSecrets(reader),
)
}
}
}
/**
* A retained epoch's decryption secrets for processing late-arriving messages.
*
@@ -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,148 @@
/*
* 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.codec.TlsReader
import com.vitorpamplona.quartz.mls.framing.MlsMessage
import com.vitorpamplona.quartz.mls.framing.PublicMessage
import com.vitorpamplona.quartz.mls.group.MlsGroup
import com.vitorpamplona.quartz.mls.group.MlsGroupState
import com.vitorpamplona.quartz.mls.schedule.KeySchedule
import kotlin.test.Test
import kotlin.test.assertContentEquals
import kotlin.test.assertEquals
import kotlin.test.assertFailsWith
import kotlin.test.assertTrue
/**
* A member that sent before seeing the latest commit sealed its message in a
* former epoch. The receiver keeps the last few epochs' receiver data and
* opens such a message with [MlsGroup.decryptFormerEpoch].
*/
class MlsGroupFormerEpochTest {
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 aLateMessageOpensFromTheFormerEpoch() {
val (alice, bob) = twoMembers()
val formerEpoch = alice.epoch
val formerExporter = alice.exporterSecret("test", ByteArray(0), 32)
val late = bob.encrypt("sent before the commit".encodeToByteArray(), "aad".encodeToByteArray())
alice.commit() // bob sent before seeing this
assertEquals(formerEpoch + 1, alice.epoch)
assertTrue(runCatching { alice.decrypt(late) }.isFailure)
val opened = alice.decryptFormerEpoch(late)
assertContentEquals("sent before the commit".encodeToByteArray(), opened.content)
assertContentEquals("aad".encodeToByteArray(), opened.authenticatedData)
assertEquals(bob.leafIndex, opened.senderLeafIndex)
assertEquals(formerEpoch, opened.epoch)
val retainedExporter = alice.formerExporterSecrets().getValue(formerEpoch)
assertContentEquals(
formerExporter,
KeySchedule.mlsExporter(retainedExporter, "test", ByteArray(0), 32),
)
}
@Test
fun aFormerEpochMessageOpensOnlyOnce() {
val (alice, bob) = twoMembers()
val late = bob.encrypt("once".encodeToByteArray())
alice.commit() // bob sent before seeing this
alice.decryptFormerEpoch(late)
assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(late) }
}
@Test
fun outOfOrderLateMessagesAllOpen() {
val (alice, bob) = twoMembers()
val first = bob.encrypt("first".encodeToByteArray())
val second = bob.encrypt("second".encodeToByteArray())
alice.commit() // bob sent before seeing this
assertContentEquals("second".encodeToByteArray(), alice.decryptFormerEpoch(second).content)
assertContentEquals("first".encodeToByteArray(), alice.decryptFormerEpoch(first).content)
}
@Test
fun theWindowSurvivesARestore() {
val (alice, bob) = twoMembers()
val first = bob.encrypt("first".encodeToByteArray())
val second = bob.encrypt("second".encodeToByteArray())
alice.commit() // bob sent before seeing this
alice.decryptFormerEpoch(second)
val restored = alice.saveAndRestore()
assertEquals(alice.retainedEpochs().map { it.epoch }, restored.retainedEpochs().map { it.epoch })
assertContentEquals("first".encodeToByteArray(), restored.decryptFormerEpoch(first).content)
// what was opened before the restore stays opened
assertTrue(runCatching { restored.decryptFormerEpoch(second) }.isFailure)
}
@Test
fun onlyTheLastEpochsAreKept() {
val (alice, bob) = twoMembers()
val tooLate = bob.encrypt("too late".encodeToByteArray())
val epochOfTooLate = alice.epoch
repeat(MlsGroup.RETAIN_EPOCHS + 1) { alice.commit() }
assertEquals(MlsGroup.RETAIN_EPOCHS, alice.retainedEpochs().size)
assertTrue(alice.retainedEpochs().none { it.epoch == epochOfTooLate })
val e = assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(tooLate) }
assertTrue(e.message!!.contains("no longer retained"))
}
@Test
fun aCurrentEpochMessageIsNotAFormerOne() {
val (alice, bob) = twoMembers()
val current = bob.encrypt("now".encodeToByteArray())
assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(current) }
assertContentEquals("now".encodeToByteArray(), alice.decrypt(current).content)
}
@Test
fun aRejectedCommitLeavesTheWindowAlone() {
val (alice, bob) = twoMembers()
val commit = bob.commit().framedCommitBytes
val pub = PublicMessage.decodeTls(TlsReader(MlsMessage.decodeTls(TlsReader(commit)).payload))
val badTag = pub.confirmationTag!!.copyOf().also { it[0] = (it[0].toInt() xor 1).toByte() }
val before = alice.retainedEpochs().map { it.epoch }
// fails on the confirmation tag, after the new epoch was derived
assertTrue(runCatching { alice.processCommit(pub.content, pub.sender.leafIndex, badTag, pub.signature) }.isFailure)
assertEquals(before, alice.retainedEpochs().map { it.epoch })
alice.processFramedCommit(commit)
assertEquals(before + (bob.epoch - 1), alice.retainedEpochs().map { it.epoch })
}
}
@@ -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). v5 appends the
// skipped-generation secrets; 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(5, version)
assertEquals(7, version)
}
@Test