From bee4130dd9b903abc52ffdf4a18392096afe1006 Mon Sep 17 00:00:00 2001 From: jeremyd <4072+jeremyd@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:55:32 -0700 Subject: [PATCH 1/2] feat(mls): open late application messages from retained former epochs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A member that sends before it has seen the latest commit seals its message in the epoch it is still in. The receiver has moved on, so decrypt() refuses it. RFC 9420 §15.2 has receivers keep recent epochs' keys for this. MlsGroup keeps the receiver data of the last RETAIN_EPOCHS (4) epochs: sender-data and encryption secrets, the sender ratchets and skipped generations, and the tree and GroupContext. It is captured before a commit changes anything, kept once the commit applies, and restored by processCommit's rollback. decryptFormerEpoch() opens an APPLICATION PrivateMessage from one of those epochs with the same checks as decrypt(), signature included, and writes the ratchet back so a generation opens only once. decrypt() is split into openPrivateMessage/applicationFromOpened so both use the same code. formerExporterSecrets() gives the retained epochs' exporter secrets. MlsGroupState v6 persists the window; older blobs decode with none. --- .../quartz/mls/group/MlsGroup.kt | 233 +++++++++++++----- .../quartz/mls/group/MlsGroupState.kt | 112 ++++++++- .../quartz/mls/MlsGroupFormerEpochTest.kt | 148 +++++++++++ .../quartz/mls/MlsGroupStateTest.kt | 6 +- 4 files changed, 435 insertions(+), 64 deletions(-) create mode 100644 quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupFormerEpochTest.kt 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 412515b562..591a9d8131 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 @@ -125,6 +125,8 @@ class MlsGroup private constructor( private var interimTranscriptHash: ByteArray, private val pskStore: MutableMap = mutableMapOf(), private val pendingProposals: MutableList = mutableListOf(), + /** Receiver data of the last [RETAIN_EPOCHS] epochs, oldest first. See [decryptFormerEpoch]. */ + private val retainedEpochs: ArrayDeque = ArrayDeque(), private val sentKeys: MutableMap = mutableMapOf(), /** Staged keys from proposeSigningKeyRotation — only promoted on successful commit */ private var pendingSigningKey: ByteArray? = null, @@ -292,6 +294,7 @@ class MlsGroup private constructor( pendingProposals = pendingProposals.toList(), skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(), skippedHandshakeSecrets = secretTree.exportSkippedHandshakeSecrets(), + retainedEpochs = retainedEpochs.toList(), ) } @@ -398,6 +401,32 @@ 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(), + ) + } + + 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 = retainedEpochs.toList() + /** * Extract retained epoch secrets for late-message decryption. * @@ -669,6 +698,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 +1058,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 +1338,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 +1363,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 +1400,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 +1424,124 @@ 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()) + + 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(), + ) + return decrypted + } + + /** Exporter secrets of the retained former epochs, by epoch, for keys an application derives per epoch. */ + fun formerExporterSecrets(): Map = 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` // 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 +1795,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 +1810,8 @@ class MlsGroup private constructor( pendingProposals.addAll(pendingSnapshot) sentKeys.clear() sentKeys.putAll(sentKeysSnapshot) + retainedEpochs.clear() + retainedEpochs.addAll(retainedSnapshot) throw t } } @@ -1716,6 +1823,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 +2217,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 +3333,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 @@ -4055,6 +4171,7 @@ class MlsGroup private constructor( interimTranscriptHash = state.interimTranscriptHash, pathPrivateKeys = state.pathPrivateKeys.toMutableMap(), pendingProposals = state.pendingProposals.toMutableList(), + retainedEpochs = ArrayDeque(state.retainedEpochs), policy = policy, ) } 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 7ea4ca5086..92a1b9e7ba 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 @@ -94,6 +94,12 @@ data class MlsGroupState( val skippedApplicationSecrets: Map, ByteArray> = emptyMap(), /** Same for the HANDSHAKE ratchet (STATE_VERSION 5+). */ val skippedHandshakeSecrets: Map, 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 = emptyList(), ) { fun encodeTls(): ByteArray { val writer = TlsWriter() @@ -171,6 +177,10 @@ 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) + return writer.toByteArray() } @@ -200,8 +210,10 @@ 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. */ - private const val STATE_VERSION = 5 + private const val STATE_VERSION = 6 fun decodeTls(data: ByteArray): MlsGroupState { val reader = TlsReader(data) @@ -306,6 +318,15 @@ 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() + } + return MlsGroupState( groupContext = groupContext, treeBytes = treeBytes, @@ -321,10 +342,11 @@ data class MlsGroupState( pendingProposals = pendingProposals, skippedApplicationSecrets = skippedApplicationSecrets, skippedHandshakeSecrets = skippedHandshakeSecrets, + retainedEpochs = retainedEpochs, ) } - private fun writeSkippedSecrets( + internal fun writeSkippedSecrets( writer: TlsWriter, secrets: Map, ByteArray>, ) { @@ -336,7 +358,7 @@ data class MlsGroupState( } } - private fun readSkippedSecrets(reader: TlsReader): Map, ByteArray> { + internal fun readSkippedSecrets(reader: TlsReader): Map, ByteArray> { val count = reader.readUint32().toInt() return buildMap { repeat(count) { @@ -349,6 +371,90 @@ 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, + val skippedApplicationSecrets: Map, ByteArray>, +) { + 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. * diff --git a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupFormerEpochTest.kt b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupFormerEpochTest.kt new file mode 100644 index 0000000000..5e6ea6bd2e --- /dev/null +++ b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupFormerEpochTest.kt @@ -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 { + 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 { 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 { alice.decryptFormerEpoch(tooLate) } + assertTrue(e.message!!.contains("no longer retained")) + } + + @Test + fun aCurrentEpochMessageIsNotAFormerOne() { + val (alice, bob) = twoMembers() + val current = bob.encrypt("now".encodeToByteArray()) + assertFailsWith { 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 }) + } +} 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 8717617f1a..1013e054ab 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). v5 appends the - // skipped-generation secrets; older blobs still decode, so the version + // First two bytes are the state version (uint16). v6 appends the + // retained former epochs; 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(6, version) } @Test From 2171b8d11d5f8af2fd3c62e3317bf1682aafc4ed Mon Sep 17 00:00:00 2001 From: jeremyd <4072+jeremyd@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:09:21 -0700 Subject: [PATCH 2/2] 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. --- .../quartz/mls/group/MlsGroup.kt | 5 + .../quartz/mls/group/MlsGroupState.kt | 50 +++++++++- .../quartz/mls/schedule/SecretTree.kt | 69 +++++++++---- .../mls/schedule/SecretTreeNodeSecretsTest.kt | 82 ++++++++++++++++ .../mls/MlsGroupStateNodeSecretsTest.kt | 96 +++++++++++++++++++ .../quartz/mls/MlsGroupStateTest.kt | 6 +- 6 files changed, 286 insertions(+), 22 deletions(-) create mode 100644 quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/mls/schedule/SecretTreeNodeSecretsTest.kt create mode 100644 quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/mls/MlsGroupStateNodeSecretsTest.kt 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