fix: VarInt encoding migration and LeafNode parent_hash for COMMIT

Major interop fixes discovered by IETF test vectors:

1. VarInt migration: All MLS TLS struct serialization now uses
   QUIC-style VarInt encoding for opaque<V> and vector<V> fields,
   matching OpenMLS and mls-rs wire format. Added readVarInt(),
   readOpaqueVarInt(), readVectorVarInt() to TlsReader and
   putVectorVarInt() to TlsWriter.

2. SecretTree left/right derivation: Fixed tree secret splitting
   to use "left"/"right" as context strings per RFC 9420 Section 9,
   instead of byte(0)/byte(1).

3. LeafNode parent_hash: Added parent_hash<V> field for COMMIT
   source per RFC 9420 Section 7.2. The COMMIT case is NOT empty -
   it includes a parent_hash opaque field.

4. MLS-Exporter: Added ByteArray overload for raw byte labels
   (test vectors use non-UTF-8 label bytes).

Test results: 32/41 passing (78%), up from 25/41 (61%).
Newly passing: SecretTree (2), TreeValidation deserialization (1),
TreeValidation resolution (1), TreeKem deserialization (1),
Commit deserialization (1), RatchetTree deserialization (1).

https://claude.ai/code/session_01NocQDWj2Y92FugjfgazzL3
This commit is contained in:
Claude
2026-04-03 20:10:40 +00:00
parent 4fc99021ae
commit 621196b74f
14 changed files with 230 additions and 122 deletions
@@ -96,6 +96,54 @@ class TlsReader(
return readBytes(length) return readBytes(length)
} }
/**
* Read a QUIC-style variable-length integer (VarInt).
* The two most significant bits of the first byte encode the length:
* - 00: 1 byte (6-bit value, 0..63)
* - 01: 2 bytes (14-bit value, 64..16383)
* - 10: 4 bytes (30-bit value, 16384..1073741823)
*/
fun readVarInt(): Int {
val first = readUint8()
return when (first shr 6) {
0 -> {
first
}
1 -> {
((first and 0x3F) shl 8) or readUint8()
}
2 -> {
val b1 = readUint8()
val b2 = readUint8()
val b3 = readUint8()
((first and 0x3F) shl 24) or (b1 shl 16) or (b2 shl 8) or b3
}
else -> {
throw IllegalArgumentException("Unsupported VarInt prefix: 0x${first.toString(16)}")
}
}
}
/** Read a variable-length opaque with QUIC-style VarInt length prefix */
fun readOpaqueVarInt(): ByteArray {
val length = readVarInt()
return readBytes(length)
}
/** Read a variable-length vector with VarInt length prefix */
fun <T> readVectorVarInt(readItem: (TlsReader) -> T): List<T> {
val vectorBytes = readOpaqueVarInt()
val vectorReader = TlsReader(vectorBytes)
val items = mutableListOf<T>()
while (vectorReader.hasRemaining) {
items.add(readItem(vectorReader))
}
return items
}
/** /**
* Read a variable-length vector with 4-byte length prefix, * Read a variable-length vector with 4-byte length prefix,
* deserializing each item using the provided factory. * deserializing each item using the provided factory.
@@ -141,6 +141,17 @@ class TlsWriter(
putOpaque4(inner.toByteArray()) putOpaque4(inner.toByteArray())
} }
/**
* Write a variable-length vector with VarInt length prefix.
*/
fun putVectorVarInt(items: List<TlsSerializable>) {
val inner = TlsWriter()
for (item in items) {
item.encodeTls(inner)
}
putOpaqueVarInt(inner.toByteArray())
}
/** /**
* Write a variable-length vector of TLS-serializable items with a 2-byte length prefix. * Write a variable-length vector of TLS-serializable items with a 2-byte length prefix.
*/ */
@@ -89,6 +89,25 @@ object MlsCryptoProvider {
return hkdfExpand(secret, hkdfLabel.toByteArray(), length) return hkdfExpand(secret, hkdfLabel.toByteArray(), length)
} }
/**
* ExpandWithLabel variant that accepts a raw byte array label.
* The "MLS 1.0 " prefix is prepended to the raw label bytes.
*/
fun expandWithLabelRaw(
secret: ByteArray,
label: ByteArray,
context: ByteArray,
length: Int,
): ByteArray {
val prefix = "MLS 1.0 ".encodeToByteArray()
val fullLabel = prefix + label
val hkdfLabel = TlsWriter(4 + fullLabel.size + context.size)
hkdfLabel.putUint16(length)
hkdfLabel.putOpaqueVarInt(fullLabel)
hkdfLabel.putOpaqueVarInt(context)
return hkdfExpand(secret, hkdfLabel.toByteArray(), length)
}
/** /**
* MLS DeriveSecret (RFC 9420 Section 8): * MLS DeriveSecret (RFC 9420 Section 8):
* DeriveSecret(Secret, Label) = ExpandWithLabel(Secret, Label, "", Nh) * DeriveSecret(Secret, Label) = ExpandWithLabel(Secret, Label, "", Nh)
@@ -268,8 +287,8 @@ data class HpkeCiphertext(
val ciphertext: ByteArray, val ciphertext: ByteArray,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque2(kemOutput) writer.putOpaqueVarInt(kemOutput)
writer.putOpaque2(ciphertext) writer.putOpaqueVarInt(ciphertext)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -287,8 +306,8 @@ data class HpkeCiphertext(
companion object { companion object {
fun decodeTls(reader: TlsReader): HpkeCiphertext = fun decodeTls(reader: TlsReader): HpkeCiphertext =
HpkeCiphertext( HpkeCiphertext(
kemOutput = reader.readOpaque2(), kemOutput = reader.readOpaqueVarInt(),
ciphertext = reader.readOpaque2(), ciphertext = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -116,10 +116,10 @@ data class FramedContentTbs(
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(version) writer.putUint16(version)
writer.putUint16(wireFormat.value) writer.putUint16(wireFormat.value)
writer.putOpaque1(groupId) writer.putOpaqueVarInt(groupId)
writer.putUint64(epoch) writer.putUint64(epoch)
encodeSender(writer, sender) encodeSender(writer, sender)
writer.putOpaque4(authenticatedData) writer.putOpaqueVarInt(authenticatedData)
writer.putUint8(contentType.value) writer.putUint8(contentType.value)
writer.putBytes(content) writer.putBytes(content)
@@ -156,22 +156,22 @@ data class PublicMessage(
val membershipTag: ByteArray? = null, val membershipTag: ByteArray? = null,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque1(groupId) writer.putOpaqueVarInt(groupId)
writer.putUint64(epoch) writer.putUint64(epoch)
encodeSender(writer, sender) encodeSender(writer, sender)
writer.putOpaque4(authenticatedData) writer.putOpaqueVarInt(authenticatedData)
writer.putUint8(contentType.value) writer.putUint8(contentType.value)
writer.putBytes(content) writer.putBytes(content)
// FramedContentAuthData // FramedContentAuthData
writer.putOpaque2(signature) writer.putOpaqueVarInt(signature)
if (contentType == ContentType.COMMIT) { if (contentType == ContentType.COMMIT) {
writer.putOpaque1(confirmationTag ?: ByteArray(0)) writer.putOpaqueVarInt(confirmationTag ?: ByteArray(0))
} }
// membership_tag (only for member senders) // membership_tag (only for member senders)
if (sender.senderType == SenderType.MEMBER) { if (sender.senderType == SenderType.MEMBER) {
writer.putOpaque1(membershipTag ?: ByteArray(0)) writer.putOpaqueVarInt(membershipTag ?: ByteArray(0))
} }
} }
@@ -185,27 +185,27 @@ data class PublicMessage(
companion object { companion object {
fun decodeTls(reader: TlsReader): PublicMessage { fun decodeTls(reader: TlsReader): PublicMessage {
val groupId = reader.readOpaque1() val groupId = reader.readOpaqueVarInt()
val epoch = reader.readUint64() val epoch = reader.readUint64()
val sender = decodeSender(reader) val sender = decodeSender(reader)
val authenticatedData = reader.readOpaque4() val authenticatedData = reader.readOpaqueVarInt()
val contentType = ContentType.fromValue(reader.readUint8()) val contentType = ContentType.fromValue(reader.readUint8())
// Content is variable based on content_type, read remaining content // Content is variable based on content_type, read remaining content
// For now, read as opaque // For now, read as opaque
val content = reader.readOpaque4() val content = reader.readOpaqueVarInt()
val signature = reader.readOpaque2() val signature = reader.readOpaqueVarInt()
val confirmationTag = val confirmationTag =
if (contentType == ContentType.COMMIT) { if (contentType == ContentType.COMMIT) {
reader.readOpaque1() reader.readOpaqueVarInt()
} else { } else {
null null
} }
val membershipTag = val membershipTag =
if (sender.senderType == SenderType.MEMBER && reader.hasRemaining) { if (sender.senderType == SenderType.MEMBER && reader.hasRemaining) {
reader.readOpaque1() reader.readOpaqueVarInt()
} else { } else {
null null
} }
@@ -239,12 +239,12 @@ data class PrivateMessage(
val ciphertext: ByteArray, val ciphertext: ByteArray,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque1(groupId) writer.putOpaqueVarInt(groupId)
writer.putUint64(epoch) writer.putUint64(epoch)
writer.putUint8(contentType.value) writer.putUint8(contentType.value)
writer.putOpaque4(authenticatedData) writer.putOpaqueVarInt(authenticatedData)
writer.putOpaque1(encryptedSenderData) writer.putOpaqueVarInt(encryptedSenderData)
writer.putOpaque4(ciphertext) writer.putOpaqueVarInt(ciphertext)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -258,12 +258,12 @@ data class PrivateMessage(
companion object { companion object {
fun decodeTls(reader: TlsReader): PrivateMessage = fun decodeTls(reader: TlsReader): PrivateMessage =
PrivateMessage( PrivateMessage(
groupId = reader.readOpaque1(), groupId = reader.readOpaqueVarInt(),
epoch = reader.readUint64(), epoch = reader.readUint64(),
contentType = ContentType.fromValue(reader.readUint8()), contentType = ContentType.fromValue(reader.readUint8()),
authenticatedData = reader.readOpaque4(), authenticatedData = reader.readOpaqueVarInt(),
encryptedSenderData = reader.readOpaque1(), encryptedSenderData = reader.readOpaqueVarInt(),
ciphertext = reader.readOpaque4(), ciphertext = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -44,14 +44,14 @@ data class Commit(
val updatePath: UpdatePath?, val updatePath: UpdatePath?,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putVector4(proposals) writer.putVectorVarInt(proposals)
writer.putOptional(updatePath) writer.putOptional(updatePath)
} }
companion object { companion object {
fun decodeTls(reader: TlsReader): Commit = fun decodeTls(reader: TlsReader): Commit =
Commit( Commit(
proposals = reader.readVector4 { ProposalOrRef.decodeTls(it) }, proposals = reader.readVectorVarInt { ProposalOrRef.decodeTls(it) },
updatePath = reader.readOptional { UpdatePath.decodeTls(it) }, updatePath = reader.readOptional { UpdatePath.decodeTls(it) },
) )
} }
@@ -77,14 +77,14 @@ data class UpdatePath(
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
leafNode.encodeTls(writer) leafNode.encodeTls(writer)
writer.putVector4(nodes) writer.putVectorVarInt(nodes)
} }
companion object { companion object {
fun decodeTls(reader: TlsReader): UpdatePath = fun decodeTls(reader: TlsReader): UpdatePath =
UpdatePath( UpdatePath(
leafNode = LeafNode.decodeTls(reader), leafNode = LeafNode.decodeTls(reader),
nodes = reader.readVector4 { UpdatePathNode.decodeTls(it) }, nodes = reader.readVectorVarInt { UpdatePathNode.decodeTls(it) },
) )
} }
} }
@@ -59,10 +59,10 @@ data class MlsKeyPackage(
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(version) writer.putUint16(version)
writer.putUint16(cipherSuite) writer.putUint16(cipherSuite)
writer.putOpaque2(initKey) writer.putOpaqueVarInt(initKey)
leafNode.encodeTls(writer) leafNode.encodeTls(writer)
writer.putVector4(extensions) writer.putVectorVarInt(extensions)
writer.putOpaque2(signature) writer.putOpaqueVarInt(signature)
} }
/** /**
@@ -81,9 +81,9 @@ data class MlsKeyPackage(
val writer = TlsWriter() val writer = TlsWriter()
writer.putUint16(version) writer.putUint16(version)
writer.putUint16(cipherSuite) writer.putUint16(cipherSuite)
writer.putOpaque2(initKey) writer.putOpaqueVarInt(initKey)
leafNode.encodeTls(writer) leafNode.encodeTls(writer)
writer.putVector4(extensions) writer.putVectorVarInt(extensions)
return writer.toByteArray() return writer.toByteArray()
} }
@@ -108,10 +108,10 @@ data class MlsKeyPackage(
MlsKeyPackage( MlsKeyPackage(
version = reader.readUint16(), version = reader.readUint16(),
cipherSuite = reader.readUint16(), cipherSuite = reader.readUint16(),
initKey = reader.readOpaque2(), initKey = reader.readOpaqueVarInt(),
leafNode = LeafNode.decodeTls(reader), leafNode = LeafNode.decodeTls(reader),
extensions = reader.readVector4 { Extension.decodeTls(it) }, extensions = reader.readVectorVarInt { Extension.decodeTls(it) },
signature = reader.readOpaque2(), signature = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -127,7 +127,7 @@ sealed class Proposal : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(proposalType.value) writer.putUint16(proposalType.value)
writer.putVector4(extensions) writer.putVectorVarInt(extensions)
} }
} }
@@ -144,8 +144,8 @@ sealed class Proposal : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(proposalType.value) writer.putUint16(proposalType.value)
writer.putUint8(pskType) writer.putUint8(pskType)
writer.putOpaque2(pskId) writer.putOpaqueVarInt(pskId)
writer.putOpaque1(pskNonce) writer.putOpaqueVarInt(pskNonce)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -178,11 +178,11 @@ sealed class Proposal : TlsSerializable {
} }
ProposalType.GROUP_CONTEXT_EXTENSIONS -> { ProposalType.GROUP_CONTEXT_EXTENSIONS -> {
GroupContextExtensions(reader.readVector4 { Extension.decodeTls(it) }) GroupContextExtensions(reader.readVectorVarInt { Extension.decodeTls(it) })
} }
ProposalType.PSK -> { ProposalType.PSK -> {
Psk(reader.readUint8(), reader.readOpaque2(), reader.readOpaque1()) Psk(reader.readUint8(), reader.readOpaqueVarInt(), reader.readOpaqueVarInt())
} }
else -> { else -> {
@@ -211,7 +211,7 @@ sealed class ProposalOrRef : TlsSerializable {
) : ProposalOrRef() { ) : ProposalOrRef() {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint8(2) // reference writer.putUint8(2) // reference
writer.putOpaque1(proposalRef) writer.putOpaqueVarInt(proposalRef)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -228,7 +228,7 @@ sealed class ProposalOrRef : TlsSerializable {
val type = reader.readUint8() val type = reader.readUint8()
return when (type) { return when (type) {
1 -> Inline(Proposal.decodeTls(reader)) 1 -> Inline(Proposal.decodeTls(reader))
2 -> Reference(reader.readOpaque1()) 2 -> Reference(reader.readOpaqueVarInt())
else -> throw IllegalArgumentException("Unknown ProposalOrRef type: $type") else -> throw IllegalArgumentException("Unknown ProposalOrRef type: $type")
} }
} }
@@ -49,8 +49,8 @@ data class Welcome(
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(cipherSuite) writer.putUint16(cipherSuite)
writer.putVector4(secrets) writer.putVectorVarInt(secrets)
writer.putOpaque4(encryptedGroupInfo) writer.putOpaqueVarInt(encryptedGroupInfo)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -65,8 +65,8 @@ data class Welcome(
fun decodeTls(reader: TlsReader): Welcome = fun decodeTls(reader: TlsReader): Welcome =
Welcome( Welcome(
cipherSuite = reader.readUint16(), cipherSuite = reader.readUint16(),
secrets = reader.readVector4 { EncryptedGroupSecrets.decodeTls(it) }, secrets = reader.readVectorVarInt { EncryptedGroupSecrets.decodeTls(it) },
encryptedGroupInfo = reader.readOpaque4(), encryptedGroupInfo = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -86,7 +86,7 @@ data class EncryptedGroupSecrets(
val encryptedGroupSecrets: HpkeCiphertext, val encryptedGroupSecrets: HpkeCiphertext,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque1(newMember) writer.putOpaqueVarInt(newMember)
encryptedGroupSecrets.encodeTls(writer) encryptedGroupSecrets.encodeTls(writer)
} }
@@ -101,7 +101,7 @@ data class EncryptedGroupSecrets(
companion object { companion object {
fun decodeTls(reader: TlsReader): EncryptedGroupSecrets = fun decodeTls(reader: TlsReader): EncryptedGroupSecrets =
EncryptedGroupSecrets( EncryptedGroupSecrets(
newMember = reader.readOpaque1(), newMember = reader.readOpaqueVarInt(),
encryptedGroupSecrets = HpkeCiphertext.decodeTls(reader), encryptedGroupSecrets = HpkeCiphertext.decodeTls(reader),
) )
} }
@@ -131,10 +131,10 @@ data class GroupInfo(
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
groupContext.encodeTls(writer) groupContext.encodeTls(writer)
writer.putVector4(extensions) writer.putVectorVarInt(extensions)
writer.putOpaque1(confirmationTag) writer.putOpaqueVarInt(confirmationTag)
writer.putUint32(signer.toLong()) writer.putUint32(signer.toLong())
writer.putOpaque2(signature) writer.putOpaqueVarInt(signature)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -149,10 +149,10 @@ data class GroupInfo(
fun decodeTls(reader: TlsReader): GroupInfo = fun decodeTls(reader: TlsReader): GroupInfo =
GroupInfo( GroupInfo(
groupContext = GroupContext.decodeTls(reader), groupContext = GroupContext.decodeTls(reader),
extensions = reader.readVector4 { Extension.decodeTls(it) }, extensions = reader.readVectorVarInt { Extension.decodeTls(it) },
confirmationTag = reader.readOpaque1(), confirmationTag = reader.readOpaqueVarInt(),
signer = reader.readUint32().toInt(), signer = reader.readUint32().toInt(),
signature = reader.readOpaque2(), signature = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -186,11 +186,11 @@ data class GroupContext(
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(version) writer.putUint16(version)
writer.putUint16(cipherSuite) writer.putUint16(cipherSuite)
writer.putOpaque1(groupId) writer.putOpaqueVarInt(groupId)
writer.putUint64(epoch) writer.putUint64(epoch)
writer.putOpaque1(treeHash) writer.putOpaqueVarInt(treeHash)
writer.putOpaque1(confirmedTranscriptHash) writer.putOpaqueVarInt(confirmedTranscriptHash)
writer.putVector4(extensions) writer.putVectorVarInt(extensions)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -210,11 +210,11 @@ data class GroupContext(
GroupContext( GroupContext(
version = reader.readUint16(), version = reader.readUint16(),
cipherSuite = reader.readUint16(), cipherSuite = reader.readUint16(),
groupId = reader.readOpaque1(), groupId = reader.readOpaqueVarInt(),
epoch = reader.readUint64(), epoch = reader.readUint64(),
treeHash = reader.readOpaque1(), treeHash = reader.readOpaqueVarInt(),
confirmedTranscriptHash = reader.readOpaque1(), confirmedTranscriptHash = reader.readOpaqueVarInt(),
extensions = reader.readVector4 { Extension.decodeTls(it) }, extensions = reader.readVectorVarInt { Extension.decodeTls(it) },
) )
} }
} }
@@ -236,18 +236,18 @@ data class GroupSecrets(
val psks: List<ByteArray> = emptyList(), val psks: List<ByteArray> = emptyList(),
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque1(joinerSecret) writer.putOpaqueVarInt(joinerSecret)
if (pathSecret != null) { if (pathSecret != null) {
writer.putUint8(1) writer.putUint8(1)
writer.putOpaque1(pathSecret) writer.putOpaqueVarInt(pathSecret)
} else { } else {
writer.putUint8(0) writer.putUint8(0)
} }
val pskWriter = TlsWriter() val pskWriter = TlsWriter()
for (psk in psks) { for (psk in psks) {
pskWriter.putOpaque2(psk) pskWriter.putOpaqueVarInt(psk)
} }
writer.putOpaque4(pskWriter.toByteArray()) writer.putOpaqueVarInt(pskWriter.toByteArray())
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -260,13 +260,13 @@ data class GroupSecrets(
companion object { companion object {
fun decodeTls(reader: TlsReader): GroupSecrets { fun decodeTls(reader: TlsReader): GroupSecrets {
val joinerSecret = reader.readOpaque1() val joinerSecret = reader.readOpaqueVarInt()
val pathSecret = reader.readOptional { it.readOpaque1() } val pathSecret = reader.readOptional { it.readOpaqueVarInt() }
val pskBytes = reader.readOpaque4() val pskBytes = reader.readOpaqueVarInt()
val pskReader = TlsReader(pskBytes) val pskReader = TlsReader(pskBytes)
val psks = mutableListOf<ByteArray>() val psks = mutableListOf<ByteArray>()
while (pskReader.hasRemaining) { while (pskReader.hasRemaining) {
psks.add(pskReader.readOpaque2()) psks.add(pskReader.readOpaqueVarInt())
} }
return GroupSecrets(joinerSecret, pathSecret, psks) return GroupSecrets(joinerSecret, pathSecret, psks)
} }
@@ -125,8 +125,18 @@ class KeySchedule(
label: String, label: String,
context: ByteArray, context: ByteArray,
length: Int, length: Int,
): ByteArray = mlsExporter(exporterSecret, label.encodeToByteArray(), context, length)
/**
* MLS-Exporter with raw byte label for non-UTF-8 labels.
*/
fun mlsExporter(
exporterSecret: ByteArray,
label: ByteArray,
context: ByteArray,
length: Int,
): ByteArray { ): ByteArray {
val derivedSecret = MlsCryptoProvider.deriveSecret(exporterSecret, label) val derivedSecret = MlsCryptoProvider.expandWithLabelRaw(exporterSecret, label, ByteArray(0), MlsCryptoProvider.HASH_OUTPUT_LENGTH)
val contextHash = MlsCryptoProvider.hash(context) val contextHash = MlsCryptoProvider.hash(context)
return MlsCryptoProvider.expandWithLabel(derivedSecret, "exported", contextHash, length) return MlsCryptoProvider.expandWithLabel(derivedSecret, "exported", contextHash, length)
} }
@@ -218,8 +218,8 @@ class SecretTree(
val leftIdx = BinaryTree.left(parentIdx) val leftIdx = BinaryTree.left(parentIdx)
val rightIdx = BinaryTree.right(parentIdx) val rightIdx = BinaryTree.right(parentIdx)
val leftSecret = MlsCryptoProvider.expandWithLabel(parentSecret, "tree", byteArrayOf(0), MlsCryptoProvider.HASH_OUTPUT_LENGTH) val leftSecret = MlsCryptoProvider.expandWithLabel(parentSecret, "tree", "left".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH)
val rightSecret = MlsCryptoProvider.expandWithLabel(parentSecret, "tree", byteArrayOf(1), MlsCryptoProvider.HASH_OUTPUT_LENGTH) val rightSecret = MlsCryptoProvider.expandWithLabel(parentSecret, "tree", "right".encodeToByteArray(), MlsCryptoProvider.HASH_OUTPUT_LENGTH)
treeSecrets[leftIdx] = leftSecret treeSecrets[leftIdx] = leftSecret
treeSecrets[rightIdx] = rightSecret treeSecrets[rightIdx] = rightSecret
@@ -34,7 +34,7 @@ sealed class Credential : TlsSerializable {
) : Credential() { ) : Credential() {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(CREDENTIAL_TYPE_BASIC) writer.putUint16(CREDENTIAL_TYPE_BASIC)
writer.putOpaque2(identity) writer.putOpaqueVarInt(identity)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -53,7 +53,7 @@ sealed class Credential : TlsSerializable {
fun decodeTls(reader: TlsReader): Credential { fun decodeTls(reader: TlsReader): Credential {
val type = reader.readUint16() val type = reader.readUint16()
return when (type) { return when (type) {
CREDENTIAL_TYPE_BASIC -> Basic(reader.readOpaque2()) CREDENTIAL_TYPE_BASIC -> Basic(reader.readOpaqueVarInt())
else -> throw IllegalArgumentException("Unknown credential type: $type") else -> throw IllegalArgumentException("Unknown credential type: $type")
} }
} }
@@ -75,27 +75,27 @@ data class Capabilities(
// versions<V><2..255> // versions<V><2..255>
val versionsWriter = TlsWriter() val versionsWriter = TlsWriter()
for (v in versions) versionsWriter.putUint16(v) for (v in versions) versionsWriter.putUint16(v)
writer.putOpaque1(versionsWriter.toByteArray()) writer.putOpaqueVarInt(versionsWriter.toByteArray())
// ciphersuites<V><2..255> // ciphersuites<V><2..255>
val csWriter = TlsWriter() val csWriter = TlsWriter()
for (cs in ciphersuites) csWriter.putUint16(cs) for (cs in ciphersuites) csWriter.putUint16(cs)
writer.putOpaque1(csWriter.toByteArray()) writer.putOpaqueVarInt(csWriter.toByteArray())
// extensions<V><2..255> // extensions<V><2..255>
val extWriter = TlsWriter() val extWriter = TlsWriter()
for (e in extensions) extWriter.putUint16(e) for (e in extensions) extWriter.putUint16(e)
writer.putOpaque1(extWriter.toByteArray()) writer.putOpaqueVarInt(extWriter.toByteArray())
// proposals<V><2..255> // proposals<V><2..255>
val propWriter = TlsWriter() val propWriter = TlsWriter()
for (p in proposals) propWriter.putUint16(p) for (p in proposals) propWriter.putUint16(p)
writer.putOpaque1(propWriter.toByteArray()) writer.putOpaqueVarInt(propWriter.toByteArray())
// credentials<V><2..255> // credentials<V><2..255>
val credWriter = TlsWriter() val credWriter = TlsWriter()
for (c in credentials) credWriter.putUint16(c) for (c in credentials) credWriter.putUint16(c)
writer.putOpaque1(credWriter.toByteArray()) writer.putOpaqueVarInt(credWriter.toByteArray())
} }
companion object { companion object {
@@ -108,11 +108,11 @@ data class Capabilities(
} }
return Capabilities( return Capabilities(
versions = readUint16List(reader.readOpaque1()), versions = readUint16List(reader.readOpaqueVarInt()),
ciphersuites = readUint16List(reader.readOpaque1()), ciphersuites = readUint16List(reader.readOpaqueVarInt()),
extensions = readUint16List(reader.readOpaque1()), extensions = readUint16List(reader.readOpaqueVarInt()),
proposals = readUint16List(reader.readOpaque1()), proposals = readUint16List(reader.readOpaqueVarInt()),
credentials = readUint16List(reader.readOpaque1()), credentials = readUint16List(reader.readOpaqueVarInt()),
) )
} }
} }
@@ -128,7 +128,7 @@ data class Extension(
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putUint16(extensionType) writer.putUint16(extensionType)
writer.putOpaque2(extensionData) writer.putOpaqueVarInt(extensionData)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -147,7 +147,7 @@ data class Extension(
fun decodeTls(reader: TlsReader): Extension = fun decodeTls(reader: TlsReader): Extension =
Extension( Extension(
extensionType = reader.readUint16(), extensionType = reader.readUint16(),
extensionData = reader.readOpaque2(), extensionData = reader.readOpaqueVarInt(),
) )
} }
} }
@@ -207,12 +207,13 @@ data class LeafNode(
val capabilities: Capabilities, val capabilities: Capabilities,
val leafNodeSource: LeafNodeSource, val leafNodeSource: LeafNodeSource,
val lifetime: Lifetime?, val lifetime: Lifetime?,
val parentHash: ByteArray? = null,
val extensions: List<Extension>, val extensions: List<Extension>,
val signature: ByteArray, val signature: ByteArray,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque2(encryptionKey) writer.putOpaqueVarInt(encryptionKey)
writer.putOpaque2(signatureKey) writer.putOpaqueVarInt(signatureKey)
writer.putStruct(credential) writer.putStruct(credential)
writer.putStruct(capabilities) writer.putStruct(capabilities)
writer.putUint8(leafNodeSource.value) writer.putUint8(leafNodeSource.value)
@@ -222,11 +223,16 @@ data class LeafNode(
lifetime?.encodeTls(writer) ?: Lifetime(0, 0).encodeTls(writer) lifetime?.encodeTls(writer) ?: Lifetime(0, 0).encodeTls(writer)
} }
LeafNodeSource.UPDATE, LeafNodeSource.COMMIT -> {} // empty LeafNodeSource.UPDATE -> {}
// empty
LeafNodeSource.COMMIT -> {
writer.putOpaqueVarInt(parentHash ?: ByteArray(0))
}
} }
writer.putVector2(extensions) writer.putVectorVarInt(extensions)
writer.putOpaque2(signature) writer.putOpaqueVarInt(signature)
} }
/** /**
@@ -239,8 +245,8 @@ data class LeafNode(
leafIndex: Int? = null, leafIndex: Int? = null,
): ByteArray { ): ByteArray {
val writer = TlsWriter() val writer = TlsWriter()
writer.putOpaque2(encryptionKey) writer.putOpaqueVarInt(encryptionKey)
writer.putOpaque2(signatureKey) writer.putOpaqueVarInt(signatureKey)
writer.putStruct(credential) writer.putStruct(credential)
writer.putStruct(capabilities) writer.putStruct(capabilities)
writer.putUint8(leafNodeSource.value) writer.putUint8(leafNodeSource.value)
@@ -250,16 +256,20 @@ data class LeafNode(
lifetime?.encodeTls(writer) ?: Lifetime(0, 0).encodeTls(writer) lifetime?.encodeTls(writer) ?: Lifetime(0, 0).encodeTls(writer)
} }
LeafNodeSource.UPDATE, LeafNodeSource.COMMIT -> {} LeafNodeSource.UPDATE -> {}
LeafNodeSource.COMMIT -> {
writer.putOpaqueVarInt(parentHash ?: ByteArray(0))
}
} }
writer.putVector2(extensions) writer.putVectorVarInt(extensions)
// Context for update/commit // Context for update/commit
if (leafNodeSource != LeafNodeSource.KEY_PACKAGE) { if (leafNodeSource != LeafNodeSource.KEY_PACKAGE) {
requireNotNull(groupId) { "group_id required for update/commit LeafNode" } requireNotNull(groupId) { "group_id required for update/commit LeafNode" }
requireNotNull(leafIndex) { "leaf_index required for update/commit LeafNode" } requireNotNull(leafIndex) { "leaf_index required for update/commit LeafNode" }
writer.putOpaque1(groupId) writer.putOpaqueVarInt(groupId)
writer.putUint32(leafIndex.toLong()) writer.putUint32(leafIndex.toLong())
} }
@@ -288,20 +298,30 @@ data class LeafNode(
companion object { companion object {
fun decodeTls(reader: TlsReader): LeafNode { fun decodeTls(reader: TlsReader): LeafNode {
val encryptionKey = reader.readOpaque2() val encryptionKey = reader.readOpaqueVarInt()
val signatureKey = reader.readOpaque2() val signatureKey = reader.readOpaqueVarInt()
val credential = Credential.decodeTls(reader) val credential = Credential.decodeTls(reader)
val capabilities = Capabilities.decodeTls(reader) val capabilities = Capabilities.decodeTls(reader)
val source = LeafNodeSource.fromValue(reader.readUint8()) val source = LeafNodeSource.fromValue(reader.readUint8())
val lifetime = var lifetime: Lifetime? = null
when (source) { var parentHash: ByteArray? = null
LeafNodeSource.KEY_PACKAGE -> Lifetime.decodeTls(reader)
else -> null when (source) {
LeafNodeSource.KEY_PACKAGE -> {
lifetime = Lifetime.decodeTls(reader)
} }
val extensions = reader.readVector2 { Extension.decodeTls(it) } LeafNodeSource.UPDATE -> {}
val signature = reader.readOpaque2()
// empty
LeafNodeSource.COMMIT -> {
parentHash = reader.readOpaqueVarInt()
}
}
val extensions = reader.readVectorVarInt { Extension.decodeTls(it) }
val signature = reader.readOpaqueVarInt()
return LeafNode( return LeafNode(
encryptionKey = encryptionKey, encryptionKey = encryptionKey,
@@ -310,6 +330,7 @@ data class LeafNode(
capabilities = capabilities, capabilities = capabilities,
leafNodeSource = source, leafNodeSource = source,
lifetime = lifetime, lifetime = lifetime,
parentHash = parentHash,
extensions = extensions, extensions = extensions,
signature = signature, signature = signature,
) )
@@ -46,14 +46,14 @@ data class ParentNode(
val unmergedLeaves: List<Int>, val unmergedLeaves: List<Int>,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque2(encryptionKey) writer.putOpaqueVarInt(encryptionKey)
writer.putOpaque1(parentHash) writer.putOpaqueVarInt(parentHash)
val ulWriter = TlsWriter() val ulWriter = TlsWriter()
for (leaf in unmergedLeaves) { for (leaf in unmergedLeaves) {
ulWriter.putUint32(leaf.toLong()) ulWriter.putUint32(leaf.toLong())
} }
writer.putOpaque4(ulWriter.toByteArray()) writer.putOpaqueVarInt(ulWriter.toByteArray())
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -73,9 +73,9 @@ data class ParentNode(
companion object { companion object {
fun decodeTls(reader: TlsReader): ParentNode { fun decodeTls(reader: TlsReader): ParentNode {
val encryptionKey = reader.readOpaque2() val encryptionKey = reader.readOpaqueVarInt()
val parentHash = reader.readOpaque1() val parentHash = reader.readOpaqueVarInt()
val ulBytes = reader.readOpaque4() val ulBytes = reader.readOpaqueVarInt()
val ulReader = TlsReader(ulBytes) val ulReader = TlsReader(ulBytes)
val unmergedLeaves = mutableListOf<Int>() val unmergedLeaves = mutableListOf<Int>()
while (ulReader.hasRemaining) { while (ulReader.hasRemaining) {
@@ -291,7 +291,7 @@ class RatchetTree(
inner.putUint8(0) inner.putUint8(0)
} }
} }
writer.putOpaque4(inner.toByteArray()) writer.putOpaqueVarInt(inner.toByteArray())
} }
private fun ensureCapacity(nodeIndex: Int) { private fun ensureCapacity(nodeIndex: Int) {
@@ -302,7 +302,7 @@ class RatchetTree(
companion object { companion object {
fun decodeTls(reader: TlsReader): RatchetTree { fun decodeTls(reader: TlsReader): RatchetTree {
val treeBytes = reader.readOpaque4() val treeBytes = reader.readOpaqueVarInt()
val treeReader = TlsReader(treeBytes) val treeReader = TlsReader(treeBytes)
val nodesList = mutableListOf<TreeNode?>() val nodesList = mutableListOf<TreeNode?>()
@@ -365,8 +365,8 @@ data class UpdatePathNode(
val encryptedPathSecret: List<com.vitorpamplona.quartz.marmot.mls.crypto.HpkeCiphertext>, val encryptedPathSecret: List<com.vitorpamplona.quartz.marmot.mls.crypto.HpkeCiphertext>,
) : TlsSerializable { ) : TlsSerializable {
override fun encodeTls(writer: TlsWriter) { override fun encodeTls(writer: TlsWriter) {
writer.putOpaque2(encryptionKey) writer.putOpaqueVarInt(encryptionKey)
writer.putVector4(encryptedPathSecret) writer.putVectorVarInt(encryptedPathSecret)
} }
override fun equals(other: Any?): Boolean { override fun equals(other: Any?): Boolean {
@@ -384,9 +384,9 @@ data class UpdatePathNode(
companion object { companion object {
fun decodeTls(reader: TlsReader): UpdatePathNode = fun decodeTls(reader: TlsReader): UpdatePathNode =
UpdatePathNode( UpdatePathNode(
encryptionKey = reader.readOpaque2(), encryptionKey = reader.readOpaqueVarInt(),
encryptedPathSecret = encryptedPathSecret =
reader.readVector4 { reader.readVectorVarInt {
com.vitorpamplona.quartz.marmot.mls.crypto.HpkeCiphertext com.vitorpamplona.quartz.marmot.mls.crypto.HpkeCiphertext
.decodeTls(it) .decodeTls(it)
}, },
@@ -139,9 +139,8 @@ class KeyScheduleInteropTest {
val secrets = ks.deriveEpochSecrets(commitSecret, initSecret, pskSecret) val secrets = ks.deriveEpochSecrets(commitSecret, initSecret, pskSecret)
// Test MLS-Exporter // Test MLS-Exporter
// The exporter label and context in the test vectors are hex-encoded byte arrays. // The exporter label in test vectors is hex-encoded raw bytes (may not be valid UTF-8).
// The mlsExporter API takes a String label, so decode hex to bytes then to String. val exporterLabel = epoch.exporter.label.hexToByteArray()
val exporterLabel = String(epoch.exporter.label.hexToByteArray())
val exporterContext = epoch.exporter.context.hexToByteArray() val exporterContext = epoch.exporter.context.hexToByteArray()
val exported = val exported =