feat: implement External Commit flow (RFC 9420 Section 8.3, 12.4.3.2)

Enables non-members to join a group without a Welcome message by using
an external commit with HPKE key encapsulation.

HPKE extensions (Hpke.kt):
- deriveKeyPair(): DHKEM(X25519) key derivation from seed
- setupBaseSExport(): HPKE sender with export-only context
- setupBaseRExport(): HPKE receiver with export-only context
- keyScheduleFull(): Full key schedule returning exporter_secret
- HpkeExportContext: Export secrets via labeled expand

MlsGroup joiner side:
- externalJoin(groupInfoBytes, identity): Static method for joining
  via external commit. Performs HPKE encapsulation to external_pub,
  derives init_secret, adds self to tree, creates Commit with
  ExternalInit proposal and UpdatePath.

MlsGroup member side:
- externalPub(): Returns the group's HPKE public key for external joins
- groupInfo(): Returns GroupInfo with ratchet_tree + external_pub extensions
- deriveExternalInitSecret(): Derives init_secret from ExternalInit kem_output
- processCommit handles ExternalInit proposals by overriding init_secret

RatchetTree fix:
- setLeaf() now expands _leafCount when setting a leaf beyond current size,
  fixing external joins where the new member extends the tree

Test: testExternalJoin verifies Alice creates group, Zara joins via external
commit, Alice processes the commit, both at epoch 1 with 2 members.

All 121 MLS tests pass (41 interop + 80 unit), 0 failures.

https://claude.ai/code/session_01NocQDWj2Y92FugjfgazzL3
This commit is contained in:
Claude
2026-04-04 00:55:09 +00:00
parent db3dadb337
commit 477d560d18
4 changed files with 346 additions and 2 deletions
@@ -39,7 +39,7 @@ object Hpke {
private val KEM_SUITE_ID = "KEM".encodeToByteArray() + byteArrayOf(0x00, 0x20) // DHKEM(X25519)
private val KDF_ID = byteArrayOf(0x00, 0x01) // HKDF-SHA256
private val AEAD_ID = byteArrayOf(0x00, 0x01) // AES-128-GCM
private val HPKE_SUITE_ID = "HPKE".encodeToByteArray() + byteArrayOf(0x00, 0x20) + KDF_ID + AEAD_ID
internal val HPKE_SUITE_ID = "HPKE".encodeToByteArray() + byteArrayOf(0x00, 0x20) + KDF_ID + AEAD_ID
private const val N_SECRET = 32 // KEM shared secret length
private const val N_ENC = 32 // KEM encapsulated key length
@@ -190,6 +190,78 @@ object Hpke {
return MlsCryptoProvider.hkdfExpand(prk, labeledInfo, length)
}
// --- DeriveKeyPair and External Init (RFC 9180 Section 7.1.3, RFC 9420 Section 8.3) ---
/**
* DeriveKeyPair for DHKEM(X25519) (RFC 9180 Section 7.1.3).
* Derives an X25519 key pair from a seed.
*/
fun deriveKeyPair(ikm: ByteArray): X25519KeyPair {
val suiteId = KEM_SUITE_ID
val dkpPrk = labeledExtract(suiteId, ByteArray(0), "dkp_prk", ikm)
val sk = labeledExpand(suiteId, dkpPrk, "sk", ByteArray(0), N_SECRET)
val pk = X25519.publicFromPrivate(sk)
return X25519KeyPair(sk, pk)
}
/**
* HPKE SetupBaseS (sender) with export-only context (RFC 9180 Section 5.1).
* Used for external init: joiner sends init_secret to group.
*
* @param recipientPub the group's external_pub
* @param info context info (empty for external init)
* @return (kem_output, exporter) where exporter derives secrets via context.export()
*/
fun setupBaseSExport(
recipientPub: ByteArray,
info: ByteArray,
): Pair<ByteArray, HpkeExportContext> {
val (sharedSecret, enc) = encap(recipientPub)
val (_, _, exporterSecret) = keyScheduleFull(sharedSecret, info)
return Pair(enc, HpkeExportContext(exporterSecret))
}
/**
* HPKE SetupBaseR (receiver) with export-only context (RFC 9180 Section 5.1).
* Used by existing members to derive init_secret from ExternalInit.
*
* @param enc the kem_output from the ExternalInit proposal
* @param recipientPriv the group's external_priv
* @param info context info (empty for external init)
* @return exporter context for deriving secrets
*/
fun setupBaseRExport(
enc: ByteArray,
recipientPriv: ByteArray,
info: ByteArray,
): HpkeExportContext {
val sharedSecret = decap(enc, recipientPriv)
val (_, _, exporterSecret) = keyScheduleFull(sharedSecret, info)
return HpkeExportContext(exporterSecret)
}
/**
* Full HPKE Key Schedule returning key, nonce, AND exporter_secret.
*/
private fun keyScheduleFull(
sharedSecret: ByteArray,
info: ByteArray,
): Triple<ByteArray, ByteArray, ByteArray> {
val mode = byteArrayOf(0x00) // Base mode
val suiteId = HPKE_SUITE_ID
val pskIdHash = labeledExtract(suiteId, ByteArray(0), "psk_id_hash", ByteArray(0))
val infoHash = labeledExtract(suiteId, ByteArray(0), "info_hash", info)
val ksContext = mode + pskIdHash + infoHash
val secret = labeledExtract(suiteId, sharedSecret, "secret", ByteArray(0))
val key = labeledExpand(suiteId, secret, "key", ksContext, N_K)
val baseNonce = labeledExpand(suiteId, secret, "base_nonce", ksContext, N_N)
val exporterSecret = labeledExpand(suiteId, secret, "exp", ksContext, N_H)
return Triple(key, baseNonce, exporterSecret)
}
// --- AEAD (AES-128-GCM) ---
private fun aeadSeal(
@@ -206,3 +278,29 @@ object Hpke {
ciphertext: ByteArray,
): ByteArray = AESGCM(key, nonce).decrypt(ciphertext, aad)
}
/**
* HPKE export-only context (RFC 9180 Section 5.3).
* Used for deriving secrets without encrypting data.
*/
class HpkeExportContext(
private val exporterSecret: ByteArray,
) {
/**
* Export a secret from the HPKE context.
* secret = LabeledExpand(exporter_secret, "sec", Hash(exporter_context), L)
*/
fun export(
exporterContext: ByteArray,
length: Int,
): ByteArray {
val suiteId = Hpke.HPKE_SUITE_ID
val contextHash = MlsCryptoProvider.hash(exporterContext)
return MlsCryptoProvider.hkdfExpand(
exporterSecret,
byteArrayOf((length shr 8).toByte(), length.toByte()) +
"HPKE-v1".encodeToByteArray() + suiteId + "sec".encodeToByteArray() + contextHash,
length,
)
}
}
@@ -23,6 +23,8 @@ package com.vitorpamplona.quartz.marmot.mls.group
import com.vitorpamplona.quartz.marmot.mls.codec.TlsReader
import com.vitorpamplona.quartz.marmot.mls.codec.TlsWriter
import com.vitorpamplona.quartz.marmot.mls.crypto.Ed25519
import com.vitorpamplona.quartz.marmot.mls.crypto.Ed25519KeyPair
import com.vitorpamplona.quartz.marmot.mls.crypto.Hpke
import com.vitorpamplona.quartz.marmot.mls.crypto.MlsCryptoProvider
import com.vitorpamplona.quartz.marmot.mls.crypto.X25519
import com.vitorpamplona.quartz.marmot.mls.framing.ContentType
@@ -622,8 +624,18 @@ class MlsGroup private constructor(
}
val pskSecret = computePskSecret(commitPskProposals)
// Check for ExternalInit proposal — overrides init_secret (RFC 9420 Section 8.3)
val externalInitProposal =
commitPskProposals.filterIsInstance<Proposal.ExternalInit>().firstOrNull()
val effectiveInitSecret =
if (externalInitProposal != null) {
deriveExternalInitSecret(externalInitProposal.kemOutput)
} else {
initSecret
}
val keySchedule = KeySchedule(groupContext.toTlsBytes())
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, initSecret, pskSecret)
epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, effectiveInitSecret, pskSecret)
initSecret = epochSecrets.initSecret
secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
@@ -661,6 +673,68 @@ class MlsGroup private constructor(
length: Int,
): ByteArray = KeySchedule.mlsExporter(epochSecrets.exporterSecret, label, context, length)
// --- External Join Support (RFC 9420 Section 8.3, 12.4.3.2) ---
/**
* Get the group's external public key for external commits.
* Non-members use this to join the group via external commit.
*
* external_priv, external_pub = KEM.DeriveKeyPair(external_secret)
*/
fun externalPub(): ByteArray {
val kp = Hpke.deriveKeyPair(epochSecrets.externalSecret)
return kp.publicKey
}
/**
* Get the GroupInfo for publishing (needed by external joiners).
* Includes the ratchet tree in extensions.
*/
fun groupInfo(): GroupInfo {
val treeWriter = TlsWriter()
tree.encodeTls(treeWriter)
val ratchetTreeExtension =
Extension(
extensionType = RATCHET_TREE_EXTENSION_TYPE,
extensionData = treeWriter.toByteArray(),
)
val externalPubExtension =
Extension(
extensionType = EXTERNAL_PUB_EXTENSION_TYPE,
extensionData = externalPub(),
)
return GroupInfo(
groupContext = groupContext,
extensions = listOf(ratchetTreeExtension, externalPubExtension),
confirmationTag =
computeConfirmationTag(
epochSecrets.confirmationKey,
groupContext.confirmedTranscriptHash,
),
signer = myLeafIndex,
signature =
MlsCryptoProvider.signWithLabel(
signingPrivateKey,
"GroupInfoTBS",
groupContext.toTlsBytes(),
),
)
}
/**
* Derive the external init_secret from an ExternalInit proposal's kem_output.
* Used by existing members when processing an external commit.
*/
internal fun deriveExternalInitSecret(kemOutput: ByteArray): ByteArray {
val externalKp = Hpke.deriveKeyPair(epochSecrets.externalSecret)
val context = Hpke.setupBaseRExport(kemOutput, externalKp.privateKey, ByteArray(0))
return context.export(
"MLS 1.0 external init secret".encodeToByteArray(),
MlsCryptoProvider.HASH_OUTPUT_LENGTH,
)
}
// --- Private Helpers ---
/**
@@ -1192,6 +1266,148 @@ class MlsGroup private constructor(
)
}
/**
* Join a group via external commit (RFC 9420 Section 12.4.3.2).
*
* The joiner obtains the GroupInfo (with ratchet tree and external_pub),
* creates an ExternalInit proposal with HPKE encapsulation, and produces
* a Commit that existing members process to add the joiner.
*
* @param groupInfoBytes TLS-serialized GroupInfo
* @param identity the joiner's identity
* @param signingKey optional Ed25519 signing key (generated if null)
* @return the new MlsGroup and the commit bytes to send to the group
*/
fun externalJoin(
groupInfoBytes: ByteArray,
identity: ByteArray,
signingKey: ByteArray? = null,
): Pair<MlsGroup, ByteArray> {
val groupInfo = GroupInfo.decodeTls(TlsReader(groupInfoBytes))
val groupContext = groupInfo.groupContext
// Extract ratchet tree from extensions
val ratchetTreeExt =
groupInfo.extensions.find { it.extensionType == RATCHET_TREE_EXTENSION_TYPE }
?: throw IllegalArgumentException("GroupInfo missing ratchet_tree extension")
val tree = RatchetTree.decodeTls(TlsReader(ratchetTreeExt.extensionData))
// Extract external_pub from extensions
val externalPubExt =
groupInfo.extensions.find { it.extensionType == EXTERNAL_PUB_EXTENSION_TYPE }
?: throw IllegalArgumentException("GroupInfo missing external_pub extension")
val externalPub = externalPubExt.extensionData
// Generate key pairs
val sigKp =
signingKey?.let { key ->
val pub = Ed25519.publicFromPrivate(key)
Ed25519KeyPair(key, pub)
} ?: Ed25519.generateKeyPair()
val encKp = X25519.generateKeyPair()
// HPKE encapsulation to external_pub to derive init_secret
val (kemOutput, exportContext) = Hpke.setupBaseSExport(externalPub, ByteArray(0))
val externalInitSecret =
exportContext.export(
"MLS 1.0 external init secret".encodeToByteArray(),
MlsCryptoProvider.HASH_OUTPUT_LENGTH,
)
// Build ExternalInit proposal
val externalInitProposal = Proposal.ExternalInit(kemOutput)
// Add ourselves to the tree
val leafNode =
buildLeafNode(
encryptionKey = encKp.publicKey,
signatureKey = sigKp.publicKey,
identity = identity,
source = LeafNodeSource.COMMIT,
signingKey = sigKp.privateKey,
groupId = groupContext.groupId,
leafIndex = tree.leafCount, // We'll be the next leaf
)
val myLeafIndex = tree.addLeaf(leafNode)
// Build UpdatePath for our leaf
val leafSecret = MlsCryptoProvider.randomBytes(MlsCryptoProvider.HASH_OUTPUT_LENGTH)
val pathSecrets = tree.derivePathSecrets(myLeafIndex, leafSecret)
val copath = BinaryTree.copath(myLeafIndex, tree.leafCount)
val pathNodes =
pathSecrets.zip(copath).map { (pathKey, copathNode) ->
val resolution = tree.resolution(copathNode)
val encryptedSecrets =
resolution.mapNotNull { resNode ->
val node = tree.getNode(resNode) ?: return@mapNotNull null
val recipientPub =
when (node) {
is com.vitorpamplona.quartz.marmot.mls.tree.TreeNode.Leaf -> {
node.leafNode.encryptionKey
}
is com.vitorpamplona.quartz.marmot.mls.tree.TreeNode.Parent -> {
node.parentNode.encryptionKey
}
}
MlsCryptoProvider.encryptWithLabel(
recipientPub,
"UpdatePathNode",
groupContext.toTlsBytes(),
pathKey.pathSecret,
)
}
UpdatePathNode(pathKey.publicKey, encryptedSecrets)
}
val updatePath = UpdatePath(leafNode, pathNodes)
tree.applyUpdatePath(myLeafIndex, updatePath.nodes)
// Build Commit
val commit =
Commit(
proposals = listOf(ProposalOrRef.Inline(externalInitProposal)),
updatePath = updatePath,
)
val commitBytes = commit.toTlsBytes()
// Derive epoch secrets using external init_secret
val commitSecret =
if (pathSecrets.isNotEmpty()) {
pathSecrets.last().pathSecret
} else {
ByteArray(MlsCryptoProvider.HASH_OUTPUT_LENGTH)
}
val newTreeHash = tree.treeHash()
val newEpoch = groupContext.epoch + 1
val newGroupContext =
groupContext.copy(
epoch = newEpoch,
treeHash = newTreeHash,
)
val keySchedule = KeySchedule(newGroupContext.toTlsBytes())
val epochSecrets = keySchedule.deriveEpochSecrets(commitSecret, externalInitSecret)
val secretTree = SecretTree(epochSecrets.encryptionSecret, tree.leafCount)
val group =
MlsGroup(
groupContext = newGroupContext,
tree = tree,
myLeafIndex = myLeafIndex,
epochSecrets = epochSecrets,
secretTree = secretTree,
initSecret = epochSecrets.initSecret,
signingPrivateKey = sigKp.privateKey,
encryptionPrivateKey = encKp.privateKey,
interimTranscriptHash = ByteArray(0),
)
return Pair(group, commitBytes)
}
/**
* Build a LeafNode with signature.
*/
@@ -64,6 +64,10 @@ class RatchetTree(
leafIndex: Int,
leafNode: LeafNode?,
) {
// Expand tree if needed
if (leafIndex >= _leafCount) {
_leafCount = leafIndex + 1
}
val nodeIdx = BinaryTree.leafToNode(leafIndex)
ensureCapacity(nodeIdx)
nodes[nodeIdx] =
@@ -243,6 +243,32 @@ class MlsGroupTest {
assertTrue(ct2.isNotEmpty())
}
@Test
fun testExternalJoin() {
// Alice creates a group
val alice = MlsGroup.create("alice".encodeToByteArray())
assertEquals(0L, alice.epoch)
assertEquals(1, alice.memberCount)
// Alice publishes GroupInfo for external joiners
val groupInfoBytes = alice.groupInfo().toTlsBytes()
// Zara joins via external commit (without a Welcome)
val (zara, commitBytes) =
MlsGroup.externalJoin(
groupInfoBytes,
"zara".encodeToByteArray(),
)
// Zara is now in the group at epoch 1
assertEquals(1L, zara.epoch)
// Alice processes Zara's external commit
alice.processCommit(commitBytes, zara.leafIndex)
assertEquals(1L, alice.epoch)
assertEquals(2, alice.memberCount)
}
@Test
fun testSelfRemove() {
val group = MlsGroup.create("alice".encodeToByteArray())