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