mirror of
https://github.com/vitorpamplona/amethyst.git
synced 2026-10-05 19:28:25 +00:00
Merge pull request #4275 from jeremyd/feat/quartz-mls-secret-tree-node-secrets
feat(mls): derive secret-tree leaves from the unexpanded node secrets
This commit is contained in:
@@ -125,6 +125,8 @@ class MlsGroup private constructor(
|
||||
private var interimTranscriptHash: ByteArray,
|
||||
private val pskStore: MutableMap<String, ByteArray> = mutableMapOf(),
|
||||
private val pendingProposals: MutableList<PendingProposal> = mutableListOf(),
|
||||
/** Receiver data of the last [RETAIN_EPOCHS] epochs, oldest first. See [decryptFormerEpoch]. */
|
||||
private val retainedEpochs: ArrayDeque<RetainedEpochReceiverData> = ArrayDeque(),
|
||||
private val sentKeys: MutableMap<Int, com.vitorpamplona.quartz.mls.schedule.KeyNonceGeneration> = mutableMapOf(),
|
||||
/** Staged keys from proposeSigningKeyRotation — only promoted on successful commit */
|
||||
private var pendingSigningKey: ByteArray? = null,
|
||||
@@ -292,6 +294,8 @@ class MlsGroup private constructor(
|
||||
pendingProposals = pendingProposals.toList(),
|
||||
skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(),
|
||||
skippedHandshakeSecrets = secretTree.exportSkippedHandshakeSecrets(),
|
||||
retainedEpochs = retainedEpochs.toList(),
|
||||
nodeSecrets = secretTree.exportNodeSecrets(),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -398,6 +402,33 @@ class MlsGroup private constructor(
|
||||
pathPrivateKeys.keys.retainAll(fullPath.toSet())
|
||||
}
|
||||
|
||||
/** The current epoch's receiver data. Taken before a commit changes anything, kept once it has applied. */
|
||||
private fun captureRetainedEpoch(): RetainedEpochReceiverData {
|
||||
val w = TlsWriter()
|
||||
tree.encodeTls(w)
|
||||
return RetainedEpochReceiverData(
|
||||
epoch = epoch,
|
||||
senderDataSecret = epochSecrets.senderDataSecret,
|
||||
encryptionSecret = epochSecrets.encryptionSecret,
|
||||
exporterSecret = epochSecrets.exporterSecret,
|
||||
resumptionPsk = epochSecrets.resumptionPsk,
|
||||
groupContext = groupContext,
|
||||
treeBytes = w.toByteArray(),
|
||||
senderRatchetStates = secretTree.exportSenderStates(),
|
||||
skippedApplicationSecrets = secretTree.exportSkippedApplicationSecrets(),
|
||||
nodeSecrets = secretTree.exportNodeSecrets(),
|
||||
)
|
||||
}
|
||||
|
||||
private fun pushRetainedEpoch(retained: RetainedEpochReceiverData) {
|
||||
retainedEpochs.removeAll { it.epoch == retained.epoch }
|
||||
retainedEpochs.addLast(retained)
|
||||
while (retainedEpochs.size > RETAIN_EPOCHS) retainedEpochs.removeFirst()
|
||||
}
|
||||
|
||||
/** Receiver data of the retained former epochs, oldest first. */
|
||||
fun retainedEpochs(): List<RetainedEpochReceiverData> = retainedEpochs.toList()
|
||||
|
||||
/**
|
||||
* Extract retained epoch secrets for late-message decryption.
|
||||
*
|
||||
@@ -669,6 +700,7 @@ class MlsGroup private constructor(
|
||||
* Returns the Commit bytes to send to the group, plus optional Welcome for new members.
|
||||
*/
|
||||
fun commit(): CommitResult {
|
||||
val retainedForThisEpoch = captureRetainedEpoch()
|
||||
val proposals = pendingProposals.toList()
|
||||
|
||||
// The application's gate on who may commit what. RFC 9420 has none of
|
||||
@@ -1028,6 +1060,7 @@ class MlsGroup private constructor(
|
||||
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, initSecret, pskSecret)
|
||||
initSecret = epochSecrets.initSecret
|
||||
secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
|
||||
pushRetainedEpoch(retainedForThisEpoch)
|
||||
|
||||
// Compute confirmation_tag and interim_transcript_hash
|
||||
val confirmationTag = computeConfirmationTag(epochSecrets.confirmationKey, newConfirmedTranscriptHash)
|
||||
@@ -1307,25 +1340,23 @@ class MlsGroup private constructor(
|
||||
null
|
||||
}
|
||||
|
||||
/** The AEAD-opened content of a PrivateMessage: who sent it and the PrivateMessageContent bytes. */
|
||||
private class OpenedPrivateMessage(
|
||||
val senderLeafIndex: Int,
|
||||
val plaintext: ByteArray,
|
||||
)
|
||||
|
||||
/**
|
||||
* Decrypt an application message from a PrivateMessage (RFC 9420 Section 6.3).
|
||||
* @throws IllegalArgumentException if the message format is invalid
|
||||
* @throws javax.crypto.AEADBadTagException if decryption fails
|
||||
* Sender-data decryption, ratchet key lookup and content AEAD for a PrivateMessage against the given
|
||||
* epoch material (RFC 9420 §6.3). Shared by [decrypt] (current epoch) and [decryptFormerEpoch].
|
||||
*/
|
||||
fun decrypt(messageBytes: ByteArray): DecryptedMessage {
|
||||
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
|
||||
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
|
||||
|
||||
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
|
||||
|
||||
// Verify epoch and group ID match current state (RFC 9420 Section 6.1)
|
||||
require(privMsg.epoch == epoch) {
|
||||
"Message epoch ${privMsg.epoch} doesn't match current epoch $epoch"
|
||||
}
|
||||
require(privMsg.groupId.contentEquals(groupId)) {
|
||||
"Message group ID doesn't match current group"
|
||||
}
|
||||
|
||||
private fun openPrivateMessage(
|
||||
privMsg: PrivateMessage,
|
||||
senderDataSecret: ByteArray,
|
||||
tree: RatchetTree,
|
||||
secretTree: SecretTree,
|
||||
allowOwnSentKeys: Boolean,
|
||||
): OpenedPrivateMessage {
|
||||
// Derive sender data key/nonce using ciphertext sample (RFC 9420 §6.3.1)
|
||||
// RFC 9420 §6.3.2: ciphertext_sample is the first KDF.Nh bytes
|
||||
// (32 for HKDF-SHA256), not AEAD.Nk (16). Using AEAD.Nk here made
|
||||
@@ -1334,14 +1365,14 @@ class MlsGroup private constructor(
|
||||
privMsg.ciphertext.copyOfRange(0, minOf(privMsg.ciphertext.size, MlsCryptoProvider.HASH_OUTPUT_LENGTH))
|
||||
val senderDataKey =
|
||||
MlsCryptoProvider.expandWithLabel(
|
||||
epochSecrets.senderDataSecret,
|
||||
senderDataSecret,
|
||||
"key",
|
||||
ciphertextSample,
|
||||
MlsCryptoProvider.AEAD_KEY_LENGTH,
|
||||
)
|
||||
val senderDataNonce =
|
||||
MlsCryptoProvider.expandWithLabel(
|
||||
epochSecrets.senderDataSecret,
|
||||
senderDataSecret,
|
||||
"nonce",
|
||||
ciphertextSample,
|
||||
MlsCryptoProvider.AEAD_NONCE_LENGTH,
|
||||
@@ -1371,7 +1402,7 @@ class MlsGroup private constructor(
|
||||
// outgoing, so every B→A commit quartz receives lands here with
|
||||
// content_type == COMMIT.
|
||||
val kng =
|
||||
if (senderLeafIndex == myLeafIndex && sentKeys.containsKey(generation)) {
|
||||
if (allowOwnSentKeys && senderLeafIndex == myLeafIndex && sentKeys.containsKey(generation)) {
|
||||
sentKeys.remove(generation)!!
|
||||
} else {
|
||||
when (privMsg.contentType) {
|
||||
@@ -1395,49 +1426,126 @@ class MlsGroup private constructor(
|
||||
val contentAad = buildPrivateContentAAD(privMsg.groupId, privMsg.epoch, privMsg.contentType, privMsg.authenticatedData)
|
||||
val pmcPlaintext = MlsCryptoProvider.aeadDecrypt(kng.key, guardedNonce, contentAad, privMsg.ciphertext)
|
||||
|
||||
return OpenedPrivateMessage(senderLeafIndex, pmcPlaintext)
|
||||
}
|
||||
|
||||
/** Parses and signature-checks an APPLICATION PrivateMessageContent against the given tree/context. */
|
||||
private fun applicationFromOpened(
|
||||
privMsg: PrivateMessage,
|
||||
opened: OpenedPrivateMessage,
|
||||
tree: RatchetTree,
|
||||
groupContext: GroupContext,
|
||||
): DecryptedMessage {
|
||||
val pmcReader = TlsReader(opened.plaintext)
|
||||
val senderLeafIndex = opened.senderLeafIndex
|
||||
val applicationData = pmcReader.readOpaqueVarInt()
|
||||
val signature = pmcReader.readOpaqueVarInt()
|
||||
while (pmcReader.hasRemaining) {
|
||||
require(pmcReader.readBytes(1)[0] == 0.toByte()) {
|
||||
"PrivateMessageContent padding must be zero"
|
||||
}
|
||||
}
|
||||
|
||||
val senderLeaf =
|
||||
requireNotNull(tree.getLeaf(senderLeafIndex)) {
|
||||
"Sender leaf is blank at index $senderLeafIndex"
|
||||
}
|
||||
require(
|
||||
MlsCryptoProvider.verifyWithLabel(
|
||||
senderLeaf.signatureKey,
|
||||
"FramedContentTBS",
|
||||
buildApplicationFramedContentTbs(
|
||||
groupId = privMsg.groupId,
|
||||
epoch = privMsg.epoch,
|
||||
senderLeafIndex = senderLeafIndex,
|
||||
authenticatedData = privMsg.authenticatedData,
|
||||
applicationData = applicationData,
|
||||
groupContext = groupContext,
|
||||
),
|
||||
signature,
|
||||
),
|
||||
) { "FramedContentTBS signature verification failed" }
|
||||
|
||||
return DecryptedMessage(
|
||||
senderLeafIndex = senderLeafIndex,
|
||||
contentType = privMsg.contentType,
|
||||
content = applicationData,
|
||||
epoch = privMsg.epoch,
|
||||
authenticatedData = privMsg.authenticatedData,
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Opens an APPLICATION message sealed in a retained former epoch: one from
|
||||
* a member that had not yet seen the latest commit (RFC 9420 §15.2 asks
|
||||
* receivers to keep recent epochs' keys for exactly this).
|
||||
*
|
||||
* Same checks as [decrypt] against that epoch's tree and GroupContext,
|
||||
* signature included, and the epoch's ratchet is written back so a
|
||||
* generation opens only once. Handshake messages from a former epoch are
|
||||
* refused, as is an epoch no longer retained.
|
||||
*/
|
||||
fun decryptFormerEpoch(messageBytes: ByteArray): DecryptedMessage {
|
||||
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
|
||||
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
|
||||
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
|
||||
require(privMsg.epoch < epoch) { "Message epoch ${privMsg.epoch} is not a former epoch (current $epoch)" }
|
||||
require(privMsg.contentType == ContentType.APPLICATION) { "Only application messages can be read from a former epoch" }
|
||||
require(privMsg.groupId.contentEquals(groupId)) { "Message group ID doesn't match current group" }
|
||||
|
||||
val index = retainedEpochs.indexOfFirst { it.epoch == privMsg.epoch }
|
||||
require(index >= 0) {
|
||||
"Message epoch ${privMsg.epoch} is no longer retained (oldest kept: ${retainedEpochs.firstOrNull()?.epoch ?: "none"})"
|
||||
}
|
||||
val retained = retainedEpochs[index]
|
||||
val formerTree = RatchetTree.decodeTls(TlsReader(retained.treeBytes))
|
||||
val formerSecrets = SecretTree(retained.encryptionSecret, formerTree.leafCount)
|
||||
formerSecrets.importSenderStates(retained.senderRatchetStates)
|
||||
formerSecrets.importSkippedSecrets(retained.skippedApplicationSecrets, emptyMap())
|
||||
if (retained.nodeSecrets.isNotEmpty()) formerSecrets.importNodeSecrets(retained.nodeSecrets)
|
||||
|
||||
val opened = openPrivateMessage(privMsg, retained.senderDataSecret, formerTree, formerSecrets, allowOwnSentKeys = false)
|
||||
val decrypted = applicationFromOpened(privMsg, opened, formerTree, retained.groupContext)
|
||||
retainedEpochs[index] =
|
||||
retained.copy(
|
||||
senderRatchetStates = formerSecrets.exportSenderStates(),
|
||||
skippedApplicationSecrets = formerSecrets.exportSkippedApplicationSecrets(),
|
||||
nodeSecrets = formerSecrets.exportNodeSecrets(),
|
||||
)
|
||||
return decrypted
|
||||
}
|
||||
|
||||
/** Exporter secrets of the retained former epochs, by epoch, for keys an application derives per epoch. */
|
||||
fun formerExporterSecrets(): Map<Long, ByteArray> = retainedEpochs.associate { it.epoch to it.exporterSecret }
|
||||
|
||||
/**
|
||||
* Decrypt an application message from a PrivateMessage (RFC 9420 Section 6.3).
|
||||
* @throws IllegalArgumentException if the message format is invalid
|
||||
* @throws javax.crypto.AEADBadTagException if decryption fails
|
||||
*/
|
||||
fun decrypt(messageBytes: ByteArray): DecryptedMessage {
|
||||
val mlsMsg = MlsMessage.decodeTls(TlsReader(messageBytes))
|
||||
require(mlsMsg.wireFormat == WireFormat.PRIVATE_MESSAGE) { "Expected PrivateMessage" }
|
||||
|
||||
val privMsg = PrivateMessage.decodeTls(TlsReader(mlsMsg.payload))
|
||||
|
||||
// Verify epoch and group ID match current state (RFC 9420 Section 6.1)
|
||||
require(privMsg.epoch == epoch) {
|
||||
"Message epoch ${privMsg.epoch} doesn't match current epoch $epoch"
|
||||
}
|
||||
require(privMsg.groupId.contentEquals(groupId)) {
|
||||
"Message group ID doesn't match current group"
|
||||
}
|
||||
|
||||
val opened = openPrivateMessage(privMsg, epochSecrets.senderDataSecret, tree, secretTree, allowOwnSentKeys = true)
|
||||
val senderLeafIndex = opened.senderLeafIndex
|
||||
// Parse PrivateMessageContent (RFC 9420 §6.3.1). The layout depends on
|
||||
// content_type — application payloads carry `opaque application_data<V>`
|
||||
// whereas commit / proposal payloads carry the struct directly (no
|
||||
// outer length prefix).
|
||||
val pmcReader = TlsReader(pmcPlaintext)
|
||||
val pmcReader = TlsReader(opened.plaintext)
|
||||
when (privMsg.contentType) {
|
||||
ContentType.APPLICATION -> {
|
||||
val applicationData = pmcReader.readOpaqueVarInt()
|
||||
val signature = pmcReader.readOpaqueVarInt()
|
||||
while (pmcReader.hasRemaining) {
|
||||
require(pmcReader.readBytes(1)[0] == 0.toByte()) {
|
||||
"PrivateMessageContent padding must be zero"
|
||||
}
|
||||
}
|
||||
|
||||
val senderLeaf =
|
||||
requireNotNull(tree.getLeaf(senderLeafIndex)) {
|
||||
"Sender leaf is blank at index $senderLeafIndex"
|
||||
}
|
||||
require(
|
||||
MlsCryptoProvider.verifyWithLabel(
|
||||
senderLeaf.signatureKey,
|
||||
"FramedContentTBS",
|
||||
buildApplicationFramedContentTbs(
|
||||
groupId = privMsg.groupId,
|
||||
epoch = privMsg.epoch,
|
||||
senderLeafIndex = senderLeafIndex,
|
||||
authenticatedData = privMsg.authenticatedData,
|
||||
applicationData = applicationData,
|
||||
groupContext = groupContext,
|
||||
),
|
||||
signature,
|
||||
),
|
||||
) { "FramedContentTBS signature verification failed" }
|
||||
|
||||
return DecryptedMessage(
|
||||
senderLeafIndex = senderLeafIndex,
|
||||
contentType = privMsg.contentType,
|
||||
content = applicationData,
|
||||
epoch = privMsg.epoch,
|
||||
authenticatedData = privMsg.authenticatedData,
|
||||
)
|
||||
}
|
||||
ContentType.APPLICATION -> return applicationFromOpened(privMsg, opened, tree, groupContext)
|
||||
|
||||
ContentType.COMMIT -> {
|
||||
// PrivateMessageContent for a Commit: the Commit struct
|
||||
@@ -1691,6 +1799,7 @@ class MlsGroup private constructor(
|
||||
val interimSnapshot = interimTranscriptHash
|
||||
val pendingSnapshot = pendingProposals.toList()
|
||||
val sentKeysSnapshot = sentKeys.toMap()
|
||||
val retainedSnapshot = retainedEpochs.toList()
|
||||
|
||||
try {
|
||||
processCommitInner(commitBytes, senderLeafIndex, confirmationTag, signature, wireFormat)
|
||||
@@ -1705,6 +1814,8 @@ class MlsGroup private constructor(
|
||||
pendingProposals.addAll(pendingSnapshot)
|
||||
sentKeys.clear()
|
||||
sentKeys.putAll(sentKeysSnapshot)
|
||||
retainedEpochs.clear()
|
||||
retainedEpochs.addAll(retainedSnapshot)
|
||||
throw t
|
||||
}
|
||||
}
|
||||
@@ -1716,6 +1827,7 @@ class MlsGroup private constructor(
|
||||
signature: ByteArray,
|
||||
wireFormat: WireFormat,
|
||||
) {
|
||||
val retainedForThisEpoch = captureRetainedEpoch()
|
||||
val commit = Commit.decodeTls(TlsReader(commitBytes))
|
||||
|
||||
// External commits (containing ExternalInit) have a sender that is not
|
||||
@@ -2109,6 +2221,7 @@ class MlsGroup private constructor(
|
||||
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, effectiveInitSecret, pskSecret)
|
||||
initSecret = epochSecrets.initSecret
|
||||
secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
|
||||
pushRetainedEpoch(retainedForThisEpoch)
|
||||
|
||||
// Verify confirmation tag (RFC 9420 Section 6.1). Every commit on
|
||||
// the wire MUST carry a confirmation_tag that matches what the
|
||||
@@ -3224,6 +3337,13 @@ class MlsGroup private constructor(
|
||||
/** MLS extensions draft `app_data_update` proposal type. */
|
||||
const val APP_DATA_UPDATE_PROPOSAL_TYPE = 0x0008
|
||||
|
||||
/**
|
||||
* How many former epochs [decryptFormerEpoch] can still read. Each one
|
||||
* keeps that epoch's decryption secrets, so this is a forward-secrecy
|
||||
* trade (RFC 9420 §15.2); 4 matches ts-mls's `retainKeysForEpochs`.
|
||||
*/
|
||||
const val RETAIN_EPOCHS = 4
|
||||
|
||||
/** How far back a fresh KeyPackage LeafNode's `not_before` is set. */
|
||||
private const val LIFETIME_SKEW_SECONDS = 3_600L
|
||||
|
||||
@@ -4042,6 +4162,7 @@ class MlsGroup private constructor(
|
||||
val secretTree = SecretTree(state.encryptionSecret, tree.leafCount)
|
||||
secretTree.importSenderStates(state.senderRatchetStates)
|
||||
secretTree.importSkippedSecrets(state.skippedApplicationSecrets, state.skippedHandshakeSecrets)
|
||||
if (state.nodeSecrets.isNotEmpty()) secretTree.importNodeSecrets(state.nodeSecrets)
|
||||
|
||||
return MlsGroup(
|
||||
groupContext = state.groupContext,
|
||||
@@ -4055,6 +4176,7 @@ class MlsGroup private constructor(
|
||||
interimTranscriptHash = state.interimTranscriptHash,
|
||||
pathPrivateKeys = state.pathPrivateKeys.toMutableMap(),
|
||||
pendingProposals = state.pendingProposals.toMutableList(),
|
||||
retainedEpochs = ArrayDeque(state.retainedEpochs),
|
||||
policy = policy,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -94,6 +94,18 @@ data class MlsGroupState(
|
||||
val skippedApplicationSecrets: Map<Pair<Int, Int>, ByteArray> = emptyMap(),
|
||||
/** Same for the HANDSHAKE ratchet (STATE_VERSION 5+). */
|
||||
val skippedHandshakeSecrets: Map<Pair<Int, Int>, ByteArray> = emptyMap(),
|
||||
/**
|
||||
* Receiver data of the last [MlsGroup.RETAIN_EPOCHS] epochs, oldest first
|
||||
* (STATE_VERSION 6+), so a late message from a former epoch still opens
|
||||
* after a restart. See [MlsGroup.decryptFormerEpoch].
|
||||
*/
|
||||
val retainedEpochs: List<RetainedEpochReceiverData> = emptyList(),
|
||||
/**
|
||||
* The secret tree's unexpanded node secrets (STATE_VERSION 7+). Empty
|
||||
* means "derive from [encryptionSecret]", which is what every older blob
|
||||
* does. A state imported from ts-mls has these instead of a root.
|
||||
*/
|
||||
val nodeSecrets: Map<Int, ByteArray> = emptyMap(),
|
||||
) {
|
||||
fun encodeTls(): ByteArray {
|
||||
val writer = TlsWriter()
|
||||
@@ -171,6 +183,15 @@ data class MlsGroupState(
|
||||
writeSkippedSecrets(writer, skippedApplicationSecrets)
|
||||
writeSkippedSecrets(writer, skippedHandshakeSecrets)
|
||||
|
||||
// Retained former epochs (STATE_VERSION 6+).
|
||||
writer.putUint32(retainedEpochs.size.toLong())
|
||||
for (retained in retainedEpochs) retained.encodeTls(writer)
|
||||
|
||||
// Secret-tree node secrets (STATE_VERSION 7+), the current epoch's,
|
||||
// then each retained epoch's in the same order as above.
|
||||
writeNodeSecrets(writer, nodeSecrets)
|
||||
for (retained in retainedEpochs) writeNodeSecrets(writer, retained.nodeSecrets)
|
||||
|
||||
return writer.toByteArray()
|
||||
}
|
||||
|
||||
@@ -200,8 +221,13 @@ data class MlsGroupState(
|
||||
* v5: appends [skippedApplicationSecrets] and [skippedHandshakeSecrets]
|
||||
* so an out-of-order message survives a restart. Older blobs
|
||||
* decode with none, as before.
|
||||
* v6: appends [retainedEpochs]. Older blobs decode with none, so the
|
||||
* first commit after the upgrade starts the window.
|
||||
* v7: appends the secret tree's [nodeSecrets], current and retained
|
||||
* epochs, so a state whose root secret is gone (imported from
|
||||
* ts-mls) restores. Older blobs derive from the root, as before.
|
||||
*/
|
||||
private const val STATE_VERSION = 5
|
||||
private const val STATE_VERSION = 7
|
||||
|
||||
fun decodeTls(data: ByteArray): MlsGroupState {
|
||||
val reader = TlsReader(data)
|
||||
@@ -306,6 +332,23 @@ data class MlsGroupState(
|
||||
val skippedApplicationSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap()
|
||||
val skippedHandshakeSecrets = if (version >= 5 && reader.hasRemaining) readSkippedSecrets(reader) else emptyMap()
|
||||
|
||||
// v6+: retained former epochs. Absent for older blobs.
|
||||
val retainedEpochs =
|
||||
if (version >= 6 && reader.hasRemaining) {
|
||||
val count = reader.readUint32().toInt()
|
||||
List(count) { RetainedEpochReceiverData.decodeTls(reader) }
|
||||
} else {
|
||||
emptyList()
|
||||
}
|
||||
|
||||
// v7+: secret-tree node secrets. Absent for older blobs.
|
||||
var nodeSecrets = emptyMap<Int, ByteArray>()
|
||||
var retainedWithNodes = retainedEpochs
|
||||
if (version >= 7 && reader.hasRemaining) {
|
||||
nodeSecrets = readNodeSecrets(reader)
|
||||
retainedWithNodes = retainedEpochs.map { it.copy(nodeSecrets = readNodeSecrets(reader)) }
|
||||
}
|
||||
|
||||
return MlsGroupState(
|
||||
groupContext = groupContext,
|
||||
treeBytes = treeBytes,
|
||||
@@ -321,10 +364,12 @@ data class MlsGroupState(
|
||||
pendingProposals = pendingProposals,
|
||||
skippedApplicationSecrets = skippedApplicationSecrets,
|
||||
skippedHandshakeSecrets = skippedHandshakeSecrets,
|
||||
retainedEpochs = retainedWithNodes,
|
||||
nodeSecrets = nodeSecrets,
|
||||
)
|
||||
}
|
||||
|
||||
private fun writeSkippedSecrets(
|
||||
internal fun writeSkippedSecrets(
|
||||
writer: TlsWriter,
|
||||
secrets: Map<Pair<Int, Int>, ByteArray>,
|
||||
) {
|
||||
@@ -336,7 +381,28 @@ data class MlsGroupState(
|
||||
}
|
||||
}
|
||||
|
||||
private fun readSkippedSecrets(reader: TlsReader): Map<Pair<Int, Int>, ByteArray> {
|
||||
private fun writeNodeSecrets(
|
||||
writer: TlsWriter,
|
||||
secrets: Map<Int, ByteArray>,
|
||||
) {
|
||||
writer.putUint32(secrets.size.toLong())
|
||||
for ((nodeIndex, secret) in secrets) {
|
||||
writer.putUint32(nodeIndex.toLong())
|
||||
writer.putOpaqueVarInt(secret)
|
||||
}
|
||||
}
|
||||
|
||||
private fun readNodeSecrets(reader: TlsReader): Map<Int, ByteArray> {
|
||||
val count = reader.readUint32().toInt()
|
||||
return buildMap {
|
||||
repeat(count) {
|
||||
val nodeIndex = reader.readUint32().toInt()
|
||||
put(nodeIndex, reader.readOpaqueVarInt())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
internal fun readSkippedSecrets(reader: TlsReader): Map<Pair<Int, Int>, ByteArray> {
|
||||
val count = reader.readUint32().toInt()
|
||||
return buildMap {
|
||||
repeat(count) {
|
||||
@@ -349,6 +415,92 @@ data class MlsGroupState(
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* What [MlsGroup.decryptFormerEpoch] needs to open a late APPLICATION message
|
||||
* from one former epoch: its sender-data and encryption secrets, the sender
|
||||
* ratchets and skipped generations as the epoch left them, and the tree and
|
||||
* GroupContext the sender's signature is checked against.
|
||||
*
|
||||
* The exporter secret and resumption PSK ride along so keys an application
|
||||
* derived in that epoch, and a resumption PSK for it (RFC 9420 §8.6), stay
|
||||
* available for as long as the epoch is retained.
|
||||
*/
|
||||
data class RetainedEpochReceiverData(
|
||||
val epoch: Long,
|
||||
val senderDataSecret: ByteArray,
|
||||
val encryptionSecret: ByteArray,
|
||||
val exporterSecret: ByteArray,
|
||||
val resumptionPsk: ByteArray,
|
||||
val groupContext: GroupContext,
|
||||
val treeBytes: ByteArray,
|
||||
val senderRatchetStates: Map<Int, SenderRatchetState>,
|
||||
val skippedApplicationSecrets: Map<Pair<Int, Int>, ByteArray>,
|
||||
/** The epoch's unexpanded secret-tree node secrets; empty means derive from [encryptionSecret]. */
|
||||
val nodeSecrets: Map<Int, ByteArray> = emptyMap(),
|
||||
) {
|
||||
override fun equals(other: Any?): Boolean {
|
||||
if (this === other) return true
|
||||
if (other !is RetainedEpochReceiverData) return false
|
||||
return epoch == other.epoch
|
||||
}
|
||||
|
||||
override fun hashCode(): Int = epoch.hashCode()
|
||||
|
||||
fun encodeTls(writer: TlsWriter) {
|
||||
writer.putUint64(epoch)
|
||||
writer.putOpaqueVarInt(senderDataSecret)
|
||||
writer.putOpaqueVarInt(encryptionSecret)
|
||||
writer.putOpaqueVarInt(exporterSecret)
|
||||
writer.putOpaqueVarInt(resumptionPsk)
|
||||
writer.putOpaqueVarInt(groupContext.toTlsBytes())
|
||||
writer.putOpaqueVarInt(treeBytes)
|
||||
writer.putUint32(senderRatchetStates.size.toLong())
|
||||
for ((leafIndex, ratchet) in senderRatchetStates) {
|
||||
writer.putUint32(leafIndex.toLong())
|
||||
writer.putOpaqueVarInt(ratchet.handshakeSecret)
|
||||
writer.putUint32(ratchet.handshakeGeneration.toLong())
|
||||
writer.putOpaqueVarInt(ratchet.applicationSecret)
|
||||
writer.putUint32(ratchet.applicationGeneration.toLong())
|
||||
}
|
||||
MlsGroupState.writeSkippedSecrets(writer, skippedApplicationSecrets)
|
||||
}
|
||||
|
||||
companion object {
|
||||
fun decodeTls(reader: TlsReader): RetainedEpochReceiverData {
|
||||
val epoch = reader.readUint64()
|
||||
val senderDataSecret = reader.readOpaqueVarInt()
|
||||
val encryptionSecret = reader.readOpaqueVarInt()
|
||||
val exporterSecret = reader.readOpaqueVarInt()
|
||||
val resumptionPsk = reader.readOpaqueVarInt()
|
||||
val groupContext = GroupContext.decodeTls(TlsReader(reader.readOpaqueVarInt()))
|
||||
val treeBytes = reader.readOpaqueVarInt()
|
||||
val count = reader.readUint32().toInt()
|
||||
val senderRatchetStates =
|
||||
buildMap {
|
||||
repeat(count) {
|
||||
val leafIndex = reader.readUint32().toInt()
|
||||
val handshakeSecret = reader.readOpaqueVarInt()
|
||||
val handshakeGeneration = reader.readUint32().toInt()
|
||||
val applicationSecret = reader.readOpaqueVarInt()
|
||||
val applicationGeneration = reader.readUint32().toInt()
|
||||
put(leafIndex, SenderRatchetState(handshakeSecret, handshakeGeneration, applicationSecret, applicationGeneration))
|
||||
}
|
||||
}
|
||||
return RetainedEpochReceiverData(
|
||||
epoch = epoch,
|
||||
senderDataSecret = senderDataSecret,
|
||||
encryptionSecret = encryptionSecret,
|
||||
exporterSecret = exporterSecret,
|
||||
resumptionPsk = resumptionPsk,
|
||||
groupContext = groupContext,
|
||||
treeBytes = treeBytes,
|
||||
senderRatchetStates = senderRatchetStates,
|
||||
skippedApplicationSecrets = MlsGroupState.readSkippedSecrets(reader),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* A retained epoch's decryption secrets for processing late-arriving messages.
|
||||
*
|
||||
|
||||
@@ -44,7 +44,7 @@ import com.vitorpamplona.quartz.mls.tree.BinaryTree
|
||||
* ```
|
||||
*/
|
||||
class SecretTree(
|
||||
private val encryptionSecret: ByteArray,
|
||||
encryptionSecret: ByteArray,
|
||||
private val leafCount: Int,
|
||||
) {
|
||||
/** Per-sender ratchet state: (handshake generation, handshake secret, app generation, app secret) */
|
||||
@@ -70,6 +70,22 @@ class SecretTree(
|
||||
/** Same cache for the HANDSHAKE ratchet. */
|
||||
private val handshakeSkippedKeys = mutableMapOf<Pair<Int, Int>, ByteArray>()
|
||||
|
||||
/**
|
||||
* Tree-node secrets not expanded yet, by node index. Starts as
|
||||
* `{root: encryptionSecret}`. Deriving a leaf replaces each node on the
|
||||
* way down with its other child, the way RFC 9420 §9 describes and ts-mls
|
||||
* keeps it (`intermediateNodes`).
|
||||
*
|
||||
* A state imported from another implementation usually no longer has the
|
||||
* root: once a sender has been derived, only the siblings along its path
|
||||
* are left. [importNodeSecrets] takes that map and later senders are
|
||||
* derived from their nearest known ancestor.
|
||||
*/
|
||||
private val nodeSecrets =
|
||||
mutableMapOf<Int, ByteArray>().also {
|
||||
if (encryptionSecret.isNotEmpty()) it[BinaryTree.root(leafCount)] = encryptionSecret
|
||||
}
|
||||
|
||||
private companion object {
|
||||
/** Maximum number of skipped key entries to retain (prevents unbounded memory growth). */
|
||||
const val MAX_SKIPPED_KEYS = 1000
|
||||
@@ -366,8 +382,9 @@ class SecretTree(
|
||||
}
|
||||
|
||||
/**
|
||||
* Derive the leaf secret from the encryption secret by walking DOWN the
|
||||
* binary tree from the root to the target leaf (RFC 9420 §9).
|
||||
* Derive the leaf secret by walking DOWN the binary tree from the nearest
|
||||
* ancestor in [nodeSecrets] (the root, in a fresh tree) to the target
|
||||
* leaf (RFC 9420 §9).
|
||||
*
|
||||
* At each step we pick left or right based on which subtree contains the
|
||||
* target. In an MLS left-balanced tree the left-subtree node indices are
|
||||
@@ -393,23 +410,32 @@ class SecretTree(
|
||||
*/
|
||||
private fun getLeafSecret(leafIndex: Int): ByteArray {
|
||||
val targetNode = BinaryTree.leafToNode(leafIndex)
|
||||
val rootIdx = BinaryTree.root(leafCount)
|
||||
var currentSecret = encryptionSecret
|
||||
var currentNode = rootIdx
|
||||
nodeSecrets.remove(targetNode)?.let { return it }
|
||||
|
||||
while (currentNode != targetNode) {
|
||||
val goLeft = targetNode < currentNode
|
||||
val label = if (goLeft) "left" else "right"
|
||||
currentSecret =
|
||||
MlsCryptoProvider.expandWithLabel(
|
||||
currentSecret,
|
||||
"tree",
|
||||
label.encodeToByteArray(),
|
||||
MlsCryptoProvider.HASH_OUTPUT_LENGTH,
|
||||
)
|
||||
currentNode = if (goLeft) BinaryTree.left(currentNode) else BinaryTree.right(currentNode)
|
||||
val path = mutableListOf(BinaryTree.root(leafCount))
|
||||
while (path.last() != targetNode) {
|
||||
val node = path.last()
|
||||
path.add(if (targetNode < node) BinaryTree.left(node) else BinaryTree.right(node))
|
||||
}
|
||||
|
||||
// Start from the nearest ancestor whose secret is still known.
|
||||
val start = path.indexOfLast { it in nodeSecrets }
|
||||
require(start >= 0) { "No secret-tree node secret left to derive leaf $leafIndex" }
|
||||
|
||||
var currentSecret = nodeSecrets.remove(path[start])!!
|
||||
for (k in start until path.size - 1) {
|
||||
val node = path[k]
|
||||
val left = MlsCryptoProvider.expandWithLabel(currentSecret, "tree", "left".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH)
|
||||
val right = MlsCryptoProvider.expandWithLabel(currentSecret, "tree", "right".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH)
|
||||
// keep the other child's secret for the senders under it
|
||||
if (path[k + 1] == BinaryTree.left(node)) {
|
||||
nodeSecrets[BinaryTree.right(node)] = right
|
||||
currentSecret = left
|
||||
} else {
|
||||
nodeSecrets[BinaryTree.left(node)] = left
|
||||
currentSecret = right
|
||||
}
|
||||
}
|
||||
return currentSecret
|
||||
}
|
||||
|
||||
@@ -451,6 +477,15 @@ class SecretTree(
|
||||
/** Same for the HANDSHAKE ratchet. */
|
||||
fun exportSkippedHandshakeSecrets(): Map<Pair<Int, Int>, ByteArray> = handshakeSkippedKeys.toMap()
|
||||
|
||||
/** Tree-node secrets not expanded yet, by node index (ts-mls `intermediateNodes`). */
|
||||
fun exportNodeSecrets(): Map<Int, ByteArray> = nodeSecrets.toMap()
|
||||
|
||||
/** Replaces the unexpanded node secrets, e.g. with a state saved by another implementation. */
|
||||
fun importNodeSecrets(secrets: Map<Int, ByteArray>) {
|
||||
nodeSecrets.clear()
|
||||
nodeSecrets.putAll(secrets)
|
||||
}
|
||||
|
||||
/** Restores skipped-generation secrets, up to the usual cache bound per ratchet. */
|
||||
fun importSkippedSecrets(
|
||||
application: Map<Pair<Int, Int>, ByteArray>,
|
||||
|
||||
+82
@@ -0,0 +1,82 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Vitor Pamplona
|
||||
*
|
||||
* Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||
* this software and associated documentation files (the "Software"), to deal in
|
||||
* the Software without restriction, including without limitation the rights to use,
|
||||
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
|
||||
* Software, and to permit persons to whom the Software is furnished to do so,
|
||||
* subject to the following conditions:
|
||||
*
|
||||
* The above copyright notice and this permission notice shall be included in all
|
||||
* copies or substantial portions of the Software.
|
||||
*
|
||||
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
|
||||
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
*/
|
||||
package com.vitorpamplona.quartz.mls.schedule
|
||||
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertContentEquals
|
||||
import kotlin.test.assertEquals
|
||||
import kotlin.test.assertFailsWith
|
||||
|
||||
class SecretTreeNodeSecretsTest {
|
||||
private val encryptionSecret = ByteArray(32) { 9 }
|
||||
|
||||
private fun key(
|
||||
tree: SecretTree,
|
||||
leafIndex: Int,
|
||||
) = tree.applicationKeyNonceForGeneration(leafIndex, 0).key
|
||||
|
||||
@Test
|
||||
fun derivingALeafKeepsTheSiblingsAndDropsThePath() {
|
||||
// 4 leaves: leaf nodes 0, 2, 4, 6; parents 1 and 5; root 3
|
||||
val tree = SecretTree(encryptionSecret, leafCount = 4)
|
||||
assertEquals(setOf(3), tree.exportNodeSecrets().keys)
|
||||
|
||||
key(tree, 0)
|
||||
assertEquals(setOf(2, 5), tree.exportNodeSecrets().keys)
|
||||
|
||||
key(tree, 3)
|
||||
assertEquals(setOf(2, 4), tree.exportNodeSecrets().keys)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aTreeWithoutItsRootDerivesTheSameKeys() {
|
||||
val reference = SecretTree(encryptionSecret, leafCount = 4)
|
||||
val source = SecretTree(encryptionSecret, leafCount = 4)
|
||||
key(source, 0)
|
||||
|
||||
// what another implementation hands over once leaf 0 has sent: no root
|
||||
val imported = SecretTree(ByteArray(0), leafCount = 4)
|
||||
imported.importNodeSecrets(source.exportNodeSecrets())
|
||||
|
||||
for (leaf in 1..3) assertContentEquals(key(reference, leaf), key(imported, leaf))
|
||||
}
|
||||
|
||||
@Test
|
||||
fun nonPowerOfTwoTreesDeriveTheSameKeys() {
|
||||
val reference = SecretTree(encryptionSecret, leafCount = 5)
|
||||
val source = SecretTree(encryptionSecret, leafCount = 5)
|
||||
key(source, 4)
|
||||
|
||||
val imported = SecretTree(ByteArray(0), leafCount = 5)
|
||||
imported.importNodeSecrets(source.exportNodeSecrets())
|
||||
for (leaf in 0..3) assertContentEquals(key(reference, leaf), key(imported, leaf))
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aLeafWithNoKnownAncestorFails() {
|
||||
val source = SecretTree(encryptionSecret, leafCount = 4)
|
||||
key(source, 0)
|
||||
val imported = SecretTree(ByteArray(0), leafCount = 4)
|
||||
imported.importNodeSecrets(source.exportNodeSecrets().filterKeys { it != 5 })
|
||||
|
||||
assertFailsWith<IllegalArgumentException> { key(imported, 2) }
|
||||
}
|
||||
}
|
||||
+148
@@ -0,0 +1,148 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Vitor Pamplona
|
||||
*
|
||||
* Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||
* this software and associated documentation files (the "Software"), to deal in
|
||||
* the Software without restriction, including without limitation the rights to use,
|
||||
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
|
||||
* Software, and to permit persons to whom the Software is furnished to do so,
|
||||
* subject to the following conditions:
|
||||
*
|
||||
* The above copyright notice and this permission notice shall be included in all
|
||||
* copies or substantial portions of the Software.
|
||||
*
|
||||
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
|
||||
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
*/
|
||||
package com.vitorpamplona.quartz.mls
|
||||
|
||||
import com.vitorpamplona.quartz.mls.codec.TlsReader
|
||||
import com.vitorpamplona.quartz.mls.framing.MlsMessage
|
||||
import com.vitorpamplona.quartz.mls.framing.PublicMessage
|
||||
import com.vitorpamplona.quartz.mls.group.MlsGroup
|
||||
import com.vitorpamplona.quartz.mls.group.MlsGroupState
|
||||
import com.vitorpamplona.quartz.mls.schedule.KeySchedule
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertContentEquals
|
||||
import kotlin.test.assertEquals
|
||||
import kotlin.test.assertFailsWith
|
||||
import kotlin.test.assertTrue
|
||||
|
||||
/**
|
||||
* A member that sent before seeing the latest commit sealed its message in a
|
||||
* former epoch. The receiver keeps the last few epochs' receiver data and
|
||||
* opens such a message with [MlsGroup.decryptFormerEpoch].
|
||||
*/
|
||||
class MlsGroupFormerEpochTest {
|
||||
private fun twoMembers(): Pair<MlsGroup, MlsGroup> {
|
||||
val alice = MlsGroup.create("alice".encodeToByteArray())
|
||||
val bobBundle =
|
||||
MlsGroup
|
||||
.create("bob".encodeToByteArray())
|
||||
.createKeyPackage("bob".encodeToByteArray(), ByteArray(0))
|
||||
val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle)
|
||||
return alice to bob
|
||||
}
|
||||
|
||||
private fun MlsGroup.saveAndRestore() = MlsGroup.restore(MlsGroupState.decodeTls(saveState().encodeTls()))
|
||||
|
||||
@Test
|
||||
fun aLateMessageOpensFromTheFormerEpoch() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val formerEpoch = alice.epoch
|
||||
val formerExporter = alice.exporterSecret("test", ByteArray(0), 32)
|
||||
val late = bob.encrypt("sent before the commit".encodeToByteArray(), "aad".encodeToByteArray())
|
||||
alice.commit() // bob sent before seeing this
|
||||
|
||||
assertEquals(formerEpoch + 1, alice.epoch)
|
||||
assertTrue(runCatching { alice.decrypt(late) }.isFailure)
|
||||
|
||||
val opened = alice.decryptFormerEpoch(late)
|
||||
assertContentEquals("sent before the commit".encodeToByteArray(), opened.content)
|
||||
assertContentEquals("aad".encodeToByteArray(), opened.authenticatedData)
|
||||
assertEquals(bob.leafIndex, opened.senderLeafIndex)
|
||||
assertEquals(formerEpoch, opened.epoch)
|
||||
|
||||
val retainedExporter = alice.formerExporterSecrets().getValue(formerEpoch)
|
||||
assertContentEquals(
|
||||
formerExporter,
|
||||
KeySchedule.mlsExporter(retainedExporter, "test", ByteArray(0), 32),
|
||||
)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aFormerEpochMessageOpensOnlyOnce() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val late = bob.encrypt("once".encodeToByteArray())
|
||||
alice.commit() // bob sent before seeing this
|
||||
|
||||
alice.decryptFormerEpoch(late)
|
||||
assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(late) }
|
||||
}
|
||||
|
||||
@Test
|
||||
fun outOfOrderLateMessagesAllOpen() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val first = bob.encrypt("first".encodeToByteArray())
|
||||
val second = bob.encrypt("second".encodeToByteArray())
|
||||
alice.commit() // bob sent before seeing this
|
||||
|
||||
assertContentEquals("second".encodeToByteArray(), alice.decryptFormerEpoch(second).content)
|
||||
assertContentEquals("first".encodeToByteArray(), alice.decryptFormerEpoch(first).content)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun theWindowSurvivesARestore() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val first = bob.encrypt("first".encodeToByteArray())
|
||||
val second = bob.encrypt("second".encodeToByteArray())
|
||||
alice.commit() // bob sent before seeing this
|
||||
alice.decryptFormerEpoch(second)
|
||||
|
||||
val restored = alice.saveAndRestore()
|
||||
assertEquals(alice.retainedEpochs().map { it.epoch }, restored.retainedEpochs().map { it.epoch })
|
||||
assertContentEquals("first".encodeToByteArray(), restored.decryptFormerEpoch(first).content)
|
||||
// what was opened before the restore stays opened
|
||||
assertTrue(runCatching { restored.decryptFormerEpoch(second) }.isFailure)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun onlyTheLastEpochsAreKept() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val tooLate = bob.encrypt("too late".encodeToByteArray())
|
||||
val epochOfTooLate = alice.epoch
|
||||
repeat(MlsGroup.RETAIN_EPOCHS + 1) { alice.commit() }
|
||||
|
||||
assertEquals(MlsGroup.RETAIN_EPOCHS, alice.retainedEpochs().size)
|
||||
assertTrue(alice.retainedEpochs().none { it.epoch == epochOfTooLate })
|
||||
val e = assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(tooLate) }
|
||||
assertTrue(e.message!!.contains("no longer retained"))
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aCurrentEpochMessageIsNotAFormerOne() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val current = bob.encrypt("now".encodeToByteArray())
|
||||
assertFailsWith<IllegalArgumentException> { alice.decryptFormerEpoch(current) }
|
||||
assertContentEquals("now".encodeToByteArray(), alice.decrypt(current).content)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aRejectedCommitLeavesTheWindowAlone() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val commit = bob.commit().framedCommitBytes
|
||||
val pub = PublicMessage.decodeTls(TlsReader(MlsMessage.decodeTls(TlsReader(commit)).payload))
|
||||
val badTag = pub.confirmationTag!!.copyOf().also { it[0] = (it[0].toInt() xor 1).toByte() }
|
||||
val before = alice.retainedEpochs().map { it.epoch }
|
||||
|
||||
// fails on the confirmation tag, after the new epoch was derived
|
||||
assertTrue(runCatching { alice.processCommit(pub.content, pub.sender.leafIndex, badTag, pub.signature) }.isFailure)
|
||||
assertEquals(before, alice.retainedEpochs().map { it.epoch })
|
||||
alice.processFramedCommit(commit)
|
||||
assertEquals(before + (bob.epoch - 1), alice.retainedEpochs().map { it.epoch })
|
||||
}
|
||||
}
|
||||
+96
@@ -0,0 +1,96 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Vitor Pamplona
|
||||
*
|
||||
* Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||
* this software and associated documentation files (the "Software"), to deal in
|
||||
* the Software without restriction, including without limitation the rights to use,
|
||||
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
|
||||
* Software, and to permit persons to whom the Software is furnished to do so,
|
||||
* subject to the following conditions:
|
||||
*
|
||||
* The above copyright notice and this permission notice shall be included in all
|
||||
* copies or substantial portions of the Software.
|
||||
*
|
||||
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
|
||||
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
*/
|
||||
package com.vitorpamplona.quartz.mls
|
||||
|
||||
import com.vitorpamplona.quartz.mls.group.MlsGroup
|
||||
import com.vitorpamplona.quartz.mls.group.MlsGroupState
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertContentEquals
|
||||
import kotlin.test.assertEquals
|
||||
import kotlin.test.assertTrue
|
||||
|
||||
/**
|
||||
* A saved state without the secret tree's root, as another implementation
|
||||
* (ts-mls) keeps it, restores from the unexpanded node secrets alone.
|
||||
*/
|
||||
class MlsGroupStateNodeSecretsTest {
|
||||
private fun twoMembers(): Pair<MlsGroup, MlsGroup> {
|
||||
val alice = MlsGroup.create("alice".encodeToByteArray())
|
||||
val bobBundle =
|
||||
MlsGroup
|
||||
.create("bob".encodeToByteArray())
|
||||
.createKeyPackage("bob".encodeToByteArray(), ByteArray(0))
|
||||
val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle)
|
||||
return alice to bob
|
||||
}
|
||||
|
||||
@Test
|
||||
fun aStateWithoutTheRootSecretRestores() {
|
||||
val (alice, bob) = twoMembers()
|
||||
// bob derives alice's leaf, so his tree no longer needs the root
|
||||
assertContentEquals("hi".encodeToByteArray(), bob.decrypt(alice.encrypt("hi".encodeToByteArray())).content)
|
||||
|
||||
val saved = bob.saveState()
|
||||
assertTrue(saved.nodeSecrets.isNotEmpty())
|
||||
val rootless = saved.copy(encryptionSecret = ByteArray(0))
|
||||
val restored = MlsGroup.restore(MlsGroupState.decodeTls(rootless.encodeTls()))
|
||||
|
||||
// bob's own leaf comes from the kept sibling, and alice can read it
|
||||
assertContentEquals("from bob".encodeToByteArray(), alice.decrypt(restored.encrypt("from bob".encodeToByteArray())).content)
|
||||
// alice's ratchet continues where it was
|
||||
assertContentEquals("again".encodeToByteArray(), restored.decrypt(alice.encrypt("again".encodeToByteArray())).content)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun retainedEpochsKeepTheirNodeSecrets() {
|
||||
val (alice, bob) = twoMembers()
|
||||
val late = bob.encrypt("late".encodeToByteArray())
|
||||
alice.decrypt(bob.encrypt("seen".encodeToByteArray()))
|
||||
alice.commit()
|
||||
|
||||
val retained = alice.retainedEpochs().last()
|
||||
assertTrue(retained.nodeSecrets.isNotEmpty())
|
||||
val restored = MlsGroup.restore(MlsGroupState.decodeTls(alice.saveState().encodeTls()))
|
||||
assertEquals(
|
||||
retained.nodeSecrets.keys,
|
||||
restored
|
||||
.retainedEpochs()
|
||||
.last()
|
||||
.nodeSecrets.keys,
|
||||
)
|
||||
assertContentEquals("late".encodeToByteArray(), restored.decryptFormerEpoch(late).content)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun version6StateStillDecodes() {
|
||||
val (_, bob) = twoMembers()
|
||||
val state = bob.saveState()
|
||||
val v7 = state.copy(nodeSecrets = emptyMap()).encodeTls()
|
||||
// with no retained epochs and no node secrets, v7 is v6 plus one empty count (uint32)
|
||||
assertTrue(state.retainedEpochs.isEmpty())
|
||||
val v6 = v7.copyOfRange(0, v7.size - 4)
|
||||
v6[0] = 0
|
||||
v6[1] = 6
|
||||
|
||||
val restored = MlsGroup.restore(MlsGroupState.decodeTls(v6))
|
||||
assertEquals(bob.epoch, restored.epoch)
|
||||
}
|
||||
}
|
||||
@@ -173,14 +173,14 @@ class MlsGroupStateTest {
|
||||
val state = group.saveState()
|
||||
val bytes = state.encodeTls()
|
||||
|
||||
// First two bytes are the state version (uint16). v5 appends the
|
||||
// skipped-generation secrets; older blobs still decode, so the version
|
||||
// First two bytes are the state version (uint16). v7 appends the
|
||||
// secret-tree node secrets; older blobs still decode, so the version
|
||||
// only ever moves forward when the layout gains a field.
|
||||
val reader =
|
||||
com.vitorpamplona.quartz.mls.codec
|
||||
.TlsReader(bytes)
|
||||
val version = reader.readUint16()
|
||||
assertEquals(5, version)
|
||||
assertEquals(7, version)
|
||||
}
|
||||
|
||||
@Test
|
||||
|
||||
Reference in New Issue
Block a user