diff --git a/amethyst/src/main/java/com/vitorpamplona/amethyst/ui/screen/loggedIn/chats/marmotGroup/MarmotGroupChatView.kt b/amethyst/src/main/java/com/vitorpamplona/amethyst/ui/screen/loggedIn/chats/marmotGroup/MarmotGroupChatView.kt index 604448416..ddd42eae1 100644 --- a/amethyst/src/main/java/com/vitorpamplona/amethyst/ui/screen/loggedIn/chats/marmotGroup/MarmotGroupChatView.kt +++ b/amethyst/src/main/java/com/vitorpamplona/amethyst/ui/screen/loggedIn/chats/marmotGroup/MarmotGroupChatView.kt @@ -31,6 +31,7 @@ import androidx.compose.material3.MaterialTheme import androidx.compose.material3.Text import androidx.compose.material3.TextFieldDefaults import androidx.compose.runtime.Composable +import androidx.compose.runtime.DisposableEffect import androidx.compose.runtime.derivedStateOf import androidx.compose.runtime.getValue import androidx.compose.runtime.remember @@ -74,6 +75,15 @@ fun MarmotGroupChatView( WatchLifecycleAndUpdateModel(feedViewModel) + val chatroom = remember(nostrGroupId) { + accountViewModel.account.marmotGroupList.getOrCreateGroup(nostrGroupId) + } + + DisposableEffect(nostrGroupId) { + chatroom.markAsRead() + onDispose { } + } + Column(Modifier.fillMaxHeight()) { Column( modifier = diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip00KeyPackages/KeyPackageRotationManager.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip00KeyPackages/KeyPackageRotationManager.kt index f2aad0abb..3f1058bde 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip00KeyPackages/KeyPackageRotationManager.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip00KeyPackages/KeyPackageRotationManager.kt @@ -63,7 +63,7 @@ class KeyPackageRotationManager { * @param dTagSlot the d-tag slot for addressable replacement * @return a [KeyPackageBundle] containing the KeyPackage and all private keys */ - fun generateKeyPackage( + suspend fun generateKeyPackage( identity: ByteArray, dTagSlot: String = KeyPackageUtils.PRIMARY_SLOT, ): KeyPackageBundle { @@ -96,7 +96,9 @@ class KeyPackageRotationManager { ) val bundle = KeyPackageBundle(keyPackage, initKp.privateKey, encKp.privateKey, sigKp.privateKey) - activeBundles[dTagSlot] = bundle + mutex.withLock { + activeBundles[dTagSlot] = bundle + } return bundle } @@ -120,24 +122,26 @@ class KeyPackageRotationManager { * The slot will be included in [pendingRotationSlots] and should be * rotated by the caller. */ - fun markConsumed(dTagSlot: String) { - activeBundles.remove(dTagSlot) - pendingRotations.add(dTagSlot) - } + suspend fun markConsumed(dTagSlot: String) = + mutex.withLock { + activeBundles.remove(dTagSlot) + pendingRotations.add(dTagSlot) + } /** * Mark a slot as consumed by looking up the KeyPackage reference. */ - fun markConsumedByRef(keyPackageRef: ByteArray) { - val entry = - activeBundles.entries.find { (_, bundle) -> - bundle.keyPackage.reference().contentEquals(keyPackageRef) + suspend fun markConsumedByRef(keyPackageRef: ByteArray) = + mutex.withLock { + val entry = + activeBundles.entries.find { (_, bundle) -> + bundle.keyPackage.reference().contentEquals(keyPackageRef) + } + if (entry != null) { + activeBundles.remove(entry.key) + pendingRotations.add(entry.key) } - if (entry != null) { - activeBundles.remove(entry.key) - pendingRotations.add(entry.key) } - } /** * Get the d-tag slots that need rotation (KeyPackage was consumed). diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip01Groups/MarmotGroupData.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip01Groups/MarmotGroupData.kt index 0adb3d2c6..afa39165b 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip01Groups/MarmotGroupData.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mip01Groups/MarmotGroupData.kt @@ -22,8 +22,10 @@ package com.vitorpamplona.quartz.marmot.mip01Groups import androidx.compose.runtime.Immutable import com.vitorpamplona.quartz.marmot.mls.codec.TlsReader +import com.vitorpamplona.quartz.marmot.mls.codec.TlsWriter import com.vitorpamplona.quartz.marmot.mls.tree.Extension import com.vitorpamplona.quartz.nip01Core.core.HexKey +import com.vitorpamplona.quartz.nip01Core.core.hexToByteArray import com.vitorpamplona.quartz.nip01Core.core.toHexKey /** 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 694ae2666..b38c29834 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 @@ -491,6 +491,15 @@ class MlsGroup private constructor( * Encrypt an application message as a PrivateMessage. */ fun encrypt(plaintext: ByteArray): ByteArray { + // Trim sentKeys if it grows too large + if (sentKeys.size > MAX_SENT_KEYS) { + val sortedKeys = sentKeys.keys.sorted() + val toRemove = sortedKeys.take(sentKeys.size - MAX_SENT_KEYS) + for (key in toRemove) { + sentKeys.remove(key) + } + } + val kng = secretTree.nextApplicationKeyNonce(myLeafIndex) sentKeys[kng.generation] = kng val ciphertext = MlsCryptoProvider.aeadEncrypt(kng.key, kng.nonce, ByteArray(0), plaintext) @@ -1190,6 +1199,7 @@ class MlsGroup private constructor( } companion object { + private const val MAX_SENT_KEYS = 10_000 private const val RATCHET_TREE_EXTENSION_TYPE = 0x0001 private const val REQUIRED_CAPABILITIES_EXTENSION_TYPE = 0x0002 private const val EXTERNAL_PUB_EXTENSION_TYPE = 0x0003 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 94c59c12b..bb2724898 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 @@ -56,6 +56,19 @@ class SecretTree( /** Consumed (sender, generation) pairs for replay detection (RFC 9420 Section 9.1) */ private val consumedGenerations = mutableMapOf>() + /** + * Cache of key/nonce pairs for skipped generations. + * Key: (leafIndex, generation) -> derived KeyNonceGeneration. + * When fast-forwarding a ratchet, intermediate generations are saved here + * so that out-of-order messages arriving later can still be decrypted. + */ + private val skippedKeys = mutableMapOf, KeyNonceGeneration>() + + /** Maximum number of skipped key entries to retain (prevents unbounded memory growth). */ + private companion object { + const val MAX_SKIPPED_KEYS = 1000 + } + init { // Seed the root treeSecrets[BinaryTree.root(leafCount)] = encryptionSecret @@ -111,11 +124,27 @@ class SecretTree( /** * Get the (key, nonce) for a specific generation, consuming secrets up to that point. * Used when decrypting an out-of-order message. + * + * When fast-forwarding past intermediate generations, their key/nonce pairs + * are cached in [skippedKeys] so that out-of-order messages arriving later + * can still be decrypted. */ fun applicationKeyNonceForGeneration( leafIndex: Int, generation: Int, ): KeyNonceGeneration { + // Check skipped keys cache first (out-of-order message for a previously skipped generation) + val cachedKey = skippedKeys.remove(Pair(leafIndex, generation)) + if (cachedKey != null) { + // Still mark as consumed for replay detection + val senderConsumed = consumedGenerations.getOrPut(leafIndex) { mutableSetOf() } + require(generation !in senderConsumed) { + "Replay detected: generation $generation from sender $leafIndex already consumed" + } + senderConsumed.add(generation) + return cachedKey + } + val state = getOrInitSender(leafIndex) require(generation >= state.applicationGeneration) { @@ -129,10 +158,16 @@ class SecretTree( } senderConsumed.add(generation) - // Fast-forward the ratchet + // Fast-forward the ratchet, caching intermediate key/nonce pairs var secret = state.applicationSecret var gen = state.applicationGeneration while (gen < generation) { + // Save the intermediate generation's key/nonce for later out-of-order retrieval + val intermediateKng = deriveKeyNonce(secret, gen) + val cacheKey = Pair(leafIndex, gen) + if (skippedKeys.size < MAX_SKIPPED_KEYS) { + skippedKeys[cacheKey] = intermediateKng + } secret = MlsCryptoProvider.expandWithLabel(secret, "secret", generationContext(gen), MlsCryptoProvider.HASH_OUTPUT_LENGTH) gen++ }