Merge pull request #2389 from vitorpamplona/claude/fix-marmot-group-persistence-6Cfb7

Fix GCM IV length in KeyStore decryption
This commit is contained in:
Vitor Pamplona
2026-04-14 22:02:55 -04:00
committed by GitHub
4 changed files with 78 additions and 63 deletions
@@ -38,6 +38,7 @@ class KeyStoreEncryption {
private const val TRANSFORMATION = "$ALGORITHM/$BLOCK_MODE/$PADDING"
private const val PURPOSE = KeyProperties.PURPOSE_ENCRYPT or KeyProperties.PURPOSE_DECRYPT
private const val KEY_ALIAS = "AMETHYST_AES_KEY"
private const val GCM_IV_LENGTH = 12
}
private val cipher = Cipher.getInstance(TRANSFORMATION)
@@ -93,9 +94,11 @@ class KeyStoreEncryption {
}
fun decrypt(bytes: ByteArray): ByteArray? {
// Extracts IV and decrypts the data
val iv = bytes.copyOfRange(0, cipher.blockSize)
val data = bytes.copyOfRange(cipher.blockSize, bytes.size)
// Extracts IV and decrypts the data. GCM mode uses a 12-byte IV,
// which is what cipher.iv returns in encrypt() — not the AES block
// size (16), which is what cipher.blockSize would return.
val iv = bytes.copyOfRange(0, GCM_IV_LENGTH)
val data = bytes.copyOfRange(GCM_IV_LENGTH, bytes.size)
cipher.init(Cipher.DECRYPT_MODE, getKey(), IvParameterSpec(iv))
return cipher.doFinal(data)
}
@@ -35,6 +35,8 @@ import com.vitorpamplona.quartz.marmot.mls.group.MlsGroupManager
import com.vitorpamplona.quartz.nip01Core.core.Event
import com.vitorpamplona.quartz.nip01Core.core.HexKey
import com.vitorpamplona.quartz.nip01Core.core.hexToByteArray
import kotlinx.coroutines.sync.Mutex
import kotlinx.coroutines.sync.withLock
import kotlin.io.encoding.Base64
import kotlin.io.encoding.ExperimentalEncodingApi
@@ -126,7 +128,7 @@ class MarmotInboundProcessor(
private val keyPackageRotationManager: KeyPackageRotationManager,
) {
private val commitTracker = CommitOrdering.EpochCommitTracker()
private val processedIdsLock = Any()
private val processedIdsMutex = Mutex()
private val processedEventIds = LinkedHashSet<String>()
companion object {
@@ -156,11 +158,12 @@ class MarmotInboundProcessor(
suspend fun processGroupEvent(groupEvent: GroupEvent): GroupEventResult {
// Deduplicate already-processed events (thread-safe)
val eventId = groupEvent.id
synchronized(processedIdsLock) {
if (eventId in processedEventIds) {
val gId = groupEvent.groupId()
return GroupEventResult.Duplicate(gId ?: "")
val alreadyProcessed =
processedIdsMutex.withLock {
eventId in processedEventIds
}
if (alreadyProcessed) {
return GroupEventResult.Duplicate(groupEvent.groupId() ?: "")
}
val groupId =
@@ -189,7 +192,7 @@ class MarmotInboundProcessor(
}
// Track ALL processed events for deduplication (including errors to prevent replay DoS)
synchronized(processedIdsLock) {
processedIdsMutex.withLock {
processedEventIds.add(eventId)
// Trim the set if it exceeds the max size
if (processedEventIds.size > MAX_PROCESSED_IDS) {
@@ -288,12 +291,12 @@ class MarmotInboundProcessor(
/**
* Get all (group, epoch) keys that have pending unresolved commits.
*/
fun pendingCommitGroupEpochs(): Set<CommitOrdering.GroupEpochKey> = commitTracker.pendingGroupEpochs()
suspend fun pendingCommitGroupEpochs(): Set<CommitOrdering.GroupEpochKey> = commitTracker.pendingGroupEpochs()
/**
* Clear all pending commit state.
*/
fun clearPendingCommits() {
suspend fun clearPendingCommits() {
commitTracker.clear()
}
@@ -20,6 +20,9 @@
*/
package com.vitorpamplona.quartz.marmot.mip03GroupMessages
import kotlinx.coroutines.sync.Mutex
import kotlinx.coroutines.sync.withLock
/**
* Deterministic commit conflict resolution for MLS over Nostr (MIP-03).
*
@@ -82,7 +85,7 @@ object CommitOrdering {
* to determine which commit wins for each (group, epoch).
*/
class EpochCommitTracker {
private val lock = Any()
private val mutex = Mutex()
private val pendingByGroupEpoch = mutableMapOf<GroupEpochKey, MutableList<GroupEvent>>()
companion object {
@@ -97,11 +100,11 @@ object CommitOrdering {
* @param epoch the MLS epoch number this commit targets
* @param commit the GroupEvent containing the commit
*/
fun addCommit(
suspend fun addCommit(
groupId: String,
epoch: Long,
commit: GroupEvent,
) = synchronized(lock) {
) = mutex.withLock {
val key = GroupEpochKey(groupId, epoch)
pendingByGroupEpoch.getOrPut(key) { mutableListOf() }.add(commit)
@@ -120,11 +123,11 @@ object CommitOrdering {
/**
* Returns pending commits for a specific group and epoch.
*/
fun pendingForEpoch(
suspend fun pendingForEpoch(
groupId: String,
epoch: Long,
): List<GroupEvent> =
synchronized(lock) {
mutex.withLock {
pendingByGroupEpoch[GroupEpochKey(groupId, epoch)]?.toList() ?: emptyList()
}
@@ -135,37 +138,38 @@ object CommitOrdering {
* @param epoch the MLS epoch to resolve
* @return the winning commit, or null if no commits exist for this (group, epoch)
*/
fun resolve(
suspend fun resolve(
groupId: String,
epoch: Long,
): GroupEvent? =
synchronized(lock) {
mutex.withLock {
selectWinner(pendingByGroupEpoch[GroupEpochKey(groupId, epoch)] ?: emptyList())
}
/**
* Clears pending commits for a (group, epoch) after it has been resolved.
*/
fun clearEpoch(
suspend fun clearEpoch(
groupId: String,
epoch: Long,
) = synchronized(lock) {
) = mutex.withLock {
pendingByGroupEpoch.remove(GroupEpochKey(groupId, epoch))
Unit
}
/**
* Returns all (group, epoch) keys that have pending commits.
*/
fun pendingGroupEpochs(): Set<GroupEpochKey> =
synchronized(lock) {
suspend fun pendingGroupEpochs(): Set<GroupEpochKey> =
mutex.withLock {
pendingByGroupEpoch.keys.toSet()
}
/**
* Clears all pending state.
*/
fun clear() =
synchronized(lock) {
suspend fun clear() =
mutex.withLock {
pendingByGroupEpoch.clear()
}
}
@@ -22,6 +22,7 @@ package com.vitorpamplona.quartz.marmot
import com.vitorpamplona.quartz.marmot.mip03GroupMessages.CommitOrdering
import com.vitorpamplona.quartz.marmot.mip03GroupMessages.GroupEvent
import kotlinx.coroutines.test.runTest
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertFalse
@@ -134,57 +135,61 @@ class CommitOrderingTest {
// ===== EpochCommitTracker =====
@Test
fun testEpochCommitTracker_Basic() {
val tracker = CommitOrdering.EpochCommitTracker()
val epoch1Commit1 = makeGroupEvent("bbb", 1000)
val epoch1Commit2 = makeGroupEvent("aaa", 1001)
fun testEpochCommitTracker_Basic() =
runTest {
val tracker = CommitOrdering.EpochCommitTracker()
val epoch1Commit1 = makeGroupEvent("bbb", 1000)
val epoch1Commit2 = makeGroupEvent("aaa", 1001)
tracker.addCommit(groupId, 1L, epoch1Commit1)
tracker.addCommit(groupId, 1L, epoch1Commit2)
tracker.addCommit(groupId, 1L, epoch1Commit1)
tracker.addCommit(groupId, 1L, epoch1Commit2)
assertEquals(2, tracker.pendingForEpoch(groupId, 1L).size)
assertEquals(0, tracker.pendingForEpoch(groupId, 2L).size)
assertEquals(2, tracker.pendingForEpoch(groupId, 1L).size)
assertEquals(0, tracker.pendingForEpoch(groupId, 2L).size)
// Resolve: epoch1Commit1 wins (earlier timestamp)
val winner = tracker.resolve(groupId, 1L)
assertEquals(epoch1Commit1, winner)
}
// Resolve: epoch1Commit1 wins (earlier timestamp)
val winner = tracker.resolve(groupId, 1L)
assertEquals(epoch1Commit1, winner)
}
@Test
fun testEpochCommitTracker_MultipleEpochs() {
val tracker = CommitOrdering.EpochCommitTracker()
val e1 = makeGroupEvent("aaa", 1000)
val e2 = makeGroupEvent("bbb", 2000)
fun testEpochCommitTracker_MultipleEpochs() =
runTest {
val tracker = CommitOrdering.EpochCommitTracker()
val e1 = makeGroupEvent("aaa", 1000)
val e2 = makeGroupEvent("bbb", 2000)
tracker.addCommit(groupId, 1L, e1)
tracker.addCommit(groupId, 2L, e2)
tracker.addCommit(groupId, 1L, e1)
tracker.addCommit(groupId, 2L, e2)
val expectedKeys =
setOf(
CommitOrdering.GroupEpochKey(groupId, 1L),
CommitOrdering.GroupEpochKey(groupId, 2L),
)
assertEquals(expectedKeys, tracker.pendingGroupEpochs())
val expectedKeys =
setOf(
CommitOrdering.GroupEpochKey(groupId, 1L),
CommitOrdering.GroupEpochKey(groupId, 2L),
)
assertEquals(expectedKeys, tracker.pendingGroupEpochs())
tracker.clearEpoch(groupId, 1L)
assertEquals(setOf(CommitOrdering.GroupEpochKey(groupId, 2L)), tracker.pendingGroupEpochs())
}
tracker.clearEpoch(groupId, 1L)
assertEquals(setOf(CommitOrdering.GroupEpochKey(groupId, 2L)), tracker.pendingGroupEpochs())
}
@Test
fun testEpochCommitTracker_ClearAll() {
val tracker = CommitOrdering.EpochCommitTracker()
tracker.addCommit(groupId, 1L, makeGroupEvent("aaa", 1000))
tracker.addCommit(groupId, 2L, makeGroupEvent("bbb", 2000))
fun testEpochCommitTracker_ClearAll() =
runTest {
val tracker = CommitOrdering.EpochCommitTracker()
tracker.addCommit(groupId, 1L, makeGroupEvent("aaa", 1000))
tracker.addCommit(groupId, 2L, makeGroupEvent("bbb", 2000))
tracker.clear()
tracker.clear()
assertTrue(tracker.pendingGroupEpochs().isEmpty())
assertNull(tracker.resolve(groupId, 1L))
}
assertTrue(tracker.pendingGroupEpochs().isEmpty())
assertNull(tracker.resolve(groupId, 1L))
}
@Test
fun testEpochCommitTracker_ResolveEmpty() {
val tracker = CommitOrdering.EpochCommitTracker()
assertNull(tracker.resolve(groupId, 999L))
}
fun testEpochCommitTracker_ResolveEmpty() =
runTest {
val tracker = CommitOrdering.EpochCommitTracker()
assertNull(tracker.resolve(groupId, 999L))
}
}