diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/codec/TlsWriter.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/codec/TlsWriter.kt index c0f099ffb..eecc7826f 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/codec/TlsWriter.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/codec/TlsWriter.kt @@ -107,6 +107,23 @@ class TlsWriter( putBytes(data) } + /** + * Write a QUIC-style variable-length integer prefixed opaque field. + * Uses the MLS/TLS VarInt encoding for the length: + * - 0..63: 1 byte (6-bit value, 2-bit prefix 00) + * - 64..16383: 2 bytes (14-bit value, 2-bit prefix 01) + * - 16384..1073741823: 4 bytes (30-bit value, 2-bit prefix 10) + */ + fun putOpaqueVarInt(data: ByteArray) { + val len = data.size + when { + len < 64 -> putUint8(len) + len < 16384 -> putUint16(len or 0x4000) + else -> putUint32((len.toLong() or 0x80000000L)) + } + putBytes(data) + } + /** Write a TLS-serializable struct */ fun putStruct(value: TlsSerializable) { value.encodeTls(this) diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/MlsCryptoProvider.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/MlsCryptoProvider.kt index 841803d1c..686d92de7 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/MlsCryptoProvider.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/MlsCryptoProvider.kt @@ -69,10 +69,11 @@ object MlsCryptoProvider { * * struct { * uint16 length; - * opaque label<7..255>; // "MLS 1.0 " + Label - * opaque context<0..2^32-1>; - * } HkdfLabel; + * opaque label = "MLS 1.0 " + Label; + * opaque context = Context; + * } KDFLabel; * ``` + * Label and context lengths use QUIC-style variable-length integer encoding. */ fun expandWithLabel( secret: ByteArray, @@ -83,8 +84,8 @@ object MlsCryptoProvider { val fullLabel = "MLS 1.0 $label".encodeToByteArray() val hkdfLabel = TlsWriter(4 + fullLabel.size + context.size) hkdfLabel.putUint16(length) - hkdfLabel.putOpaque1(fullLabel) - hkdfLabel.putOpaque4(context) + hkdfLabel.putOpaqueVarInt(fullLabel) + hkdfLabel.putOpaqueVarInt(context) return hkdfExpand(secret, hkdfLabel.toByteArray(), length) } @@ -106,9 +107,9 @@ object MlsCryptoProvider { value: ByteArray, ): ByteArray { val labelBytes = label.encodeToByteArray() - val writer = TlsWriter(2 + labelBytes.size + 2 + value.size) - writer.putOpaque2(labelBytes) - writer.putOpaque2(value) + val writer = TlsWriter(2 + labelBytes.size + value.size) + writer.putOpaqueVarInt(labelBytes) + writer.putOpaqueVarInt(value) return hash(writer.toByteArray()) } @@ -204,8 +205,8 @@ object MlsCryptoProvider { * MLS SignContent structure: * ``` * struct { - * opaque label<7..255>; // "MLS 1.0 " + Label - * opaque content<0..2^32-1>; + * opaque label = "MLS 1.0 " + Label; + * opaque content = Content; * } SignContent; * ``` */ @@ -215,8 +216,8 @@ object MlsCryptoProvider { ): ByteArray { val fullLabel = "MLS 1.0 $label".encodeToByteArray() val writer = TlsWriter(2 + fullLabel.size + 4 + content.size) - writer.putOpaque1(fullLabel) - writer.putOpaque4(content) + writer.putOpaqueVarInt(fullLabel) + writer.putOpaqueVarInt(content) return writer.toByteArray() } @@ -234,8 +235,8 @@ object MlsCryptoProvider { ): HpkeCiphertext { val fullLabel = "MLS 1.0 $label".encodeToByteArray() val info = TlsWriter() - info.putOpaque1(fullLabel) - info.putOpaque4(context) + info.putOpaqueVarInt(fullLabel) + info.putOpaqueVarInt(context) return Hpke.seal(publicKey, info.toByteArray(), ByteArray(0), plaintext) } @@ -252,8 +253,8 @@ object MlsCryptoProvider { ): ByteArray { val fullLabel = "MLS 1.0 $label".encodeToByteArray() val info = TlsWriter() - info.putOpaque1(fullLabel) - info.putOpaque4(context) + info.putOpaqueVarInt(fullLabel) + info.putOpaqueVarInt(context) return Hpke.open(privateKey, kemOutput, info.toByteArray(), ByteArray(0), ciphertext) } } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/schedule/SecretTree.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/schedule/SecretTree.kt index 50fa58c94..6af818609 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/schedule/SecretTree.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/schedule/SecretTree.kt @@ -65,12 +65,12 @@ class SecretTree( val state = getOrInitSender(leafIndex) val result = deriveKeyNonce(state.applicationSecret, state.applicationGeneration) - // Advance ratchet + // Advance ratchet: DeriveTreeSecret(secret_[N], "secret", N, Nh) val nextSecret = MlsCryptoProvider.expandWithLabel( state.applicationSecret, "secret", - ByteArray(0), + generationContext(state.applicationGeneration), MlsCryptoProvider.HASH_OUTPUT_LENGTH, ) senderState[leafIndex] = @@ -93,7 +93,7 @@ class SecretTree( MlsCryptoProvider.expandWithLabel( state.handshakeSecret, "secret", - ByteArray(0), + generationContext(state.handshakeGeneration), MlsCryptoProvider.HASH_OUTPUT_LENGTH, ) senderState[leafIndex] = @@ -123,14 +123,14 @@ class SecretTree( var secret = state.applicationSecret var gen = state.applicationGeneration while (gen < generation) { - secret = MlsCryptoProvider.expandWithLabel(secret, "secret", ByteArray(0), MlsCryptoProvider.HASH_OUTPUT_LENGTH) + secret = MlsCryptoProvider.expandWithLabel(secret, "secret", generationContext(gen), MlsCryptoProvider.HASH_OUTPUT_LENGTH) gen++ } val result = deriveKeyNonce(secret, generation) // Advance past this generation - val nextSecret = MlsCryptoProvider.expandWithLabel(secret, "secret", ByteArray(0), MlsCryptoProvider.HASH_OUTPUT_LENGTH) + val nextSecret = MlsCryptoProvider.expandWithLabel(secret, "secret", generationContext(generation), MlsCryptoProvider.HASH_OUTPUT_LENGTH) senderState[leafIndex] = state.copy( applicationSecret = nextSecret, @@ -140,22 +140,38 @@ class SecretTree( return result } + /** + * Encode a generation counter as a 4-byte big-endian uint32 for DeriveTreeSecret context. + */ + private fun generationContext(generation: Int): ByteArray = + byteArrayOf( + (generation shr 24).toByte(), + (generation shr 16).toByte(), + (generation shr 8).toByte(), + generation.toByte(), + ) + + /** + * DeriveTreeSecret(Secret, Label, Generation, Length) = + * ExpandWithLabel(Secret, Label, uint32(Generation), Length) + */ private fun deriveKeyNonce( secret: ByteArray, generation: Int, ): KeyNonceGeneration { + val genCtx = generationContext(generation) val key = MlsCryptoProvider.expandWithLabel( secret, "key", - ByteArray(0), + genCtx, MlsCryptoProvider.AEAD_KEY_LENGTH, ) val nonce = MlsCryptoProvider.expandWithLabel( secret, "nonce", - ByteArray(0), + genCtx, MlsCryptoProvider.AEAD_NONCE_LENGTH, ) return KeyNonceGeneration(key, nonce, generation) diff --git a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/KeyScheduleInteropTest.kt b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/KeyScheduleInteropTest.kt index 3140aae3a..30b7371d2 100644 --- a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/KeyScheduleInteropTest.kt +++ b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/KeyScheduleInteropTest.kt @@ -21,7 +21,6 @@ package com.vitorpamplona.quartz.marmot.mls.interop import com.vitorpamplona.quartz.TestResourceLoader -import com.vitorpamplona.quartz.marmot.mls.crypto.MlsCryptoProvider import com.vitorpamplona.quartz.marmot.mls.schedule.KeySchedule import com.vitorpamplona.quartz.nip01Core.core.JsonMapper import com.vitorpamplona.quartz.nip01Core.core.hexToByteArray @@ -51,8 +50,8 @@ class KeyScheduleInteropTest { assertTrue(vectors.isNotEmpty(), "No cipher_suite==1 key-schedule vectors found") for (v in vectors) { - // For the first epoch, init_secret is all zeros - var initSecret = ByteArray(MlsCryptoProvider.HASH_OUTPUT_LENGTH) + // Use the initial_init_secret from the test vector + var initSecret = v.initialInitSecret.hexToByteArray() for ((epochIdx, epoch) in v.epochs.withIndex()) { val groupContext = epoch.groupContext.hexToByteArray() @@ -129,7 +128,7 @@ class KeyScheduleInteropTest { assertTrue(vectors.isNotEmpty(), "No cipher_suite==1 key-schedule vectors found") for (v in vectors) { - var initSecret = ByteArray(MlsCryptoProvider.HASH_OUTPUT_LENGTH) + var initSecret = v.initialInitSecret.hexToByteArray() for ((epochIdx, epoch) in v.epochs.withIndex()) { val groupContext = epoch.groupContext.hexToByteArray() diff --git a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/MlsInteropVectors.kt b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/MlsInteropVectors.kt index c0d635051..ccc4b0964 100644 --- a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/MlsInteropVectors.kt +++ b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/marmot/mls/interop/MlsInteropVectors.kt @@ -106,6 +106,8 @@ data class TreeMathVector( @Serializable data class KeyScheduleVector( @SerialName("cipher_suite") val cipherSuite: Int, + @SerialName("group_id") val groupId: String, + @SerialName("initial_init_secret") val initialInitSecret: String, val epochs: List, )