diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/group/MlsGroup.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/group/MlsGroup.kt index 3e290dc63a..a0a3dc8ca4 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/group/MlsGroup.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/group/MlsGroup.kt @@ -3638,7 +3638,9 @@ class MlsGroup private constructor( * the MIP-era set; a current-profile group passes * [buildCurrentProfileRequiredCapabilitiesExtension]. */ - requiredCapabilities: Extension = buildMarmotRequiredCapabilitiesExtension(), + requiredCapabilities: Extension? = buildMarmotRequiredCapabilitiesExtension(), + /** The group's MLS `group_id`. Null picks 32 random bytes. */ + groupId: ByteArray? = null, ): MlsGroup { val sigKp = signingKey?.let { key -> @@ -3647,7 +3649,7 @@ class MlsGroup private constructor( } ?: Ed25519.generateKeyPair() val encKp = X25519.generateKeyPair() - val groupId = MlsCryptoProvider.randomBytes(32) + val groupId = groupId ?: MlsCryptoProvider.randomBytes(32) val leafNode = buildLeafNode( @@ -3668,7 +3670,7 @@ class MlsGroup private constructor( // bake into epoch 0 (e.g. the MIP-01 MarmotGroupData extension so // new peers who join later can see the group name without first // decrypting a pre-membership bootstrap commit — see MIP-03). - val baseExtensions = listOf(requiredCapabilities) + val baseExtensions = listOfNotNull(requiredCapabilities) val groupContext = GroupContext( groupId = groupId, diff --git a/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/marmot/mls/group/GroupCreateOptionsTest.kt b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/marmot/mls/group/GroupCreateOptionsTest.kt new file mode 100644 index 0000000000..4e71e5f0c4 --- /dev/null +++ b/quartz/src/jvmAndroidTest/kotlin/com/vitorpamplona/quartz/marmot/mls/group/GroupCreateOptionsTest.kt @@ -0,0 +1,53 @@ +/* + * 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.marmot.mls.group + +import com.vitorpamplona.quartz.nip01Core.core.hexToByteArray +import kotlin.test.Test +import kotlin.test.assertContentEquals +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +class GroupCreateOptionsTest { + private val creator = "11".repeat(32).hexToByteArray() + + @Test + fun aCallerChosenGroupIdIsUsed() { + val groupId = "my-group".encodeToByteArray() + val alice = MlsGroup.create(creator, groupId = groupId) + assertContentEquals(groupId, alice.groupId) + + val bobBundle = alice.createKeyPackage("22".repeat(32).hexToByteArray(), ByteArray(0)) + val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle) + assertContentEquals(groupId, bob.groupId) + } + + @Test + fun aGroupWithoutRequiredCapabilitiesWorks() { + val alice = MlsGroup.create(creator, requiredCapabilities = null) + assertTrue(alice.groupContextExtensionsSnapshot().none { it.extensionType == 0x0003 }) + + val bobBundle = alice.createKeyPackage("22".repeat(32).hexToByteArray(), ByteArray(0)) + val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle) + assertEquals(alice.epoch, bob.epoch) + assertContentEquals("hi".encodeToByteArray(), bob.decrypt(alice.encrypt("hi".encodeToByteArray())).content) + } +}