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..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 @@ -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,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 = 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 = 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 +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, ) } 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..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 @@ -94,6 +94,18 @@ 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(), + /** + * 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() @@ -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() + 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, ByteArray>, ) { @@ -336,7 +381,28 @@ data class MlsGroupState( } } - private fun readSkippedSecrets(reader: TlsReader): Map, ByteArray> { + 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 { 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, + 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 + 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/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/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/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 8717617f1a..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). 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