mirror of
https://github.com/vitorpamplona/amethyst.git
synced 2026-10-06 03:38:23 +00:00
fix(marmot): preserve SecretTree ratchet across restore to stop generation reuse
MlsGroupState reconstructed the SecretTree from encryption_secret alone, so every restore rewound each sender's generation counter to 0. The restored local member then re-emitted generation 0 within the same epoch — reusing the AEAD key+nonce (a confidentiality break) and getting rejected by strict receivers (openmls / MDK / Whitenoise) that forbid generation reuse, per RFC 9420 §9. Two parts: - Persist per-sender ratchet positions. SecretTree gains export/importSenderStates; MlsGroupState carries them as an optional field (STATE_VERSION 2, v1 blobs still decode as empty = legacy behavior); saveState/restore wire them through. - Persist after every send. MlsGroupManager.encrypt now saves group state, not just commits — application sends advance the ratchet but previously never hit the store, so a restart between two commits still reset it. Regression tests: a peer that consumed generation 0 accepts the restored sender's next message (single + multi-send), encrypt persists the ratchet between commits, and a v1 blob still decodes/restores. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01G6uT4xzjty1xosZBkb3sHA
This commit is contained in:
@@ -215,6 +215,10 @@ class MlsGroup private constructor(
|
||||
encryptionPrivateKey = encryptionPrivateKey,
|
||||
interimTranscriptHash = interimTranscriptHash,
|
||||
encryptionSecret = epochSecrets.encryptionSecret,
|
||||
// Preserve the SecretTree ratchet positions so a restore doesn't
|
||||
// rewind our own generation counter to 0 and reuse an AEAD
|
||||
// key+nonce within this epoch (RFC 9420 §9).
|
||||
senderRatchetStates = secretTree.exportSenderStates(),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -3501,15 +3505,23 @@ class MlsGroup private constructor(
|
||||
/**
|
||||
* Restore a group from a previously saved [MlsGroupState].
|
||||
*
|
||||
* The SecretTree is reconstructed from the stored encryption_secret.
|
||||
* Note: SecretTree ratchet state (per-sender generation counters) is
|
||||
* NOT preserved — messages sent/received before the save point cannot
|
||||
* be re-decrypted, which is acceptable because they would already
|
||||
* have been processed.
|
||||
* The SecretTree is reconstructed from the stored encryption_secret,
|
||||
* then seeded with the persisted per-sender ratchet positions
|
||||
* ([MlsGroupState.senderRatchetStates]). Seeding is what keeps the
|
||||
* local member's generation counter monotonic across a restart — a
|
||||
* fresh SecretTree would restart every sender at generation 0, so our
|
||||
* next send would reuse generation 0's AEAD key+nonce within the same
|
||||
* epoch and be rejected by strict receivers (openmls / MDK /
|
||||
* Whitenoise) that forbid generation reuse.
|
||||
*
|
||||
* Receive-only ratchets that weren't persisted (STATE_VERSION 1 blobs,
|
||||
* or senders we never decrypted) simply re-derive from generation 0 on
|
||||
* first use — safe, because those messages were already processed.
|
||||
*/
|
||||
fun restore(state: MlsGroupState): MlsGroup {
|
||||
val tree = RatchetTree.decodeTls(TlsReader(state.treeBytes))
|
||||
val secretTree = SecretTree(state.encryptionSecret, tree.leafCount)
|
||||
secretTree.importSenderStates(state.senderRatchetStates)
|
||||
|
||||
return MlsGroup(
|
||||
groupContext = state.groupContext,
|
||||
|
||||
+11
-1
@@ -348,13 +348,23 @@ class MlsGroupManager(
|
||||
/**
|
||||
* Encrypt an application message.
|
||||
* Synchronized to prevent nonce reuse from concurrent encryption.
|
||||
*
|
||||
* The group state is persisted after every send. Encrypting advances the
|
||||
* SecretTree ratchet (RFC 9420 §9) but does not change the epoch, so
|
||||
* without this save a restart between two messages would reload the
|
||||
* pre-send ratchet position and re-emit an already-used generation —
|
||||
* reusing the AEAD key+nonce and getting rejected by strict receivers.
|
||||
* State was previously persisted only at commits, which left every
|
||||
* inter-commit send unprotected.
|
||||
*/
|
||||
suspend fun encrypt(
|
||||
nostrGroupId: HexKey,
|
||||
plaintext: ByteArray,
|
||||
): ByteArray =
|
||||
mutex.withLock {
|
||||
requireGroup(nostrGroupId).encrypt(plaintext)
|
||||
val ciphertext = requireGroup(nostrGroupId).encrypt(plaintext)
|
||||
persistGroup(nostrGroupId)
|
||||
ciphertext
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
+56
-2
@@ -24,6 +24,7 @@ import com.vitorpamplona.quartz.marmot.mls.codec.TlsReader
|
||||
import com.vitorpamplona.quartz.marmot.mls.codec.TlsWriter
|
||||
import com.vitorpamplona.quartz.marmot.mls.messages.GroupContext
|
||||
import com.vitorpamplona.quartz.marmot.mls.schedule.EpochSecrets
|
||||
import com.vitorpamplona.quartz.marmot.mls.schedule.SenderRatchetState
|
||||
|
||||
/**
|
||||
* Serializable snapshot of an MLS group's complete state.
|
||||
@@ -40,6 +41,13 @@ import com.vitorpamplona.quartz.marmot.mls.schedule.EpochSecrets
|
||||
*
|
||||
* Security: This blob contains secret key material (signing key, encryption key,
|
||||
* epoch secrets). It MUST be stored in encrypted local storage.
|
||||
*
|
||||
* [senderRatchetStates] carries each sender's live SecretTree ratchet position
|
||||
* (RFC 9420 §9). Preserving it is what stops the restored local member from
|
||||
* re-emitting an already-used generation within the same epoch — see
|
||||
* [com.vitorpamplona.quartz.marmot.mls.schedule.SecretTree.exportSenderStates].
|
||||
* It is optional (empty for STATE_VERSION 1 blobs) so older persisted state
|
||||
* still decodes.
|
||||
*/
|
||||
data class MlsGroupState(
|
||||
val groupContext: GroupContext,
|
||||
@@ -51,6 +59,7 @@ data class MlsGroupState(
|
||||
val encryptionPrivateKey: ByteArray,
|
||||
val interimTranscriptHash: ByteArray,
|
||||
val encryptionSecret: ByteArray,
|
||||
val senderRatchetStates: Map<Int, SenderRatchetState> = emptyMap(),
|
||||
) {
|
||||
fun encodeTls(): ByteArray {
|
||||
val writer = TlsWriter()
|
||||
@@ -94,6 +103,18 @@ data class MlsGroupState(
|
||||
// Encryption secret for SecretTree reconstruction
|
||||
writer.putOpaqueVarInt(encryptionSecret)
|
||||
|
||||
// Per-sender SecretTree ratchet positions (STATE_VERSION 2+).
|
||||
// Preserving the local sender's generation counter is what prevents
|
||||
// AEAD key+nonce reuse (and strict-receiver rejection) after a restore.
|
||||
writer.putUint32(senderRatchetStates.size.toLong())
|
||||
for ((leafIndex, ratchet) in senderRatchetStates) {
|
||||
writer.putUint32(leafIndex.toLong())
|
||||
writer.putOpaqueVarInt(ratchet.handshakeSecret)
|
||||
writer.putUint32(ratchet.handshakeGeneration.toLong())
|
||||
writer.putOpaqueVarInt(ratchet.applicationSecret)
|
||||
writer.putUint32(ratchet.applicationGeneration.toLong())
|
||||
}
|
||||
|
||||
return writer.toByteArray()
|
||||
}
|
||||
|
||||
@@ -110,13 +131,18 @@ data class MlsGroupState(
|
||||
}
|
||||
|
||||
companion object {
|
||||
private const val STATE_VERSION = 1
|
||||
/**
|
||||
* v1: original layout (no SecretTree ratchet positions).
|
||||
* v2: appends [senderRatchetStates] so restores don't reset the
|
||||
* ratchet to generation 0. v1 blobs still decode (empty map).
|
||||
*/
|
||||
private const val STATE_VERSION = 2
|
||||
|
||||
fun decodeTls(data: ByteArray): MlsGroupState {
|
||||
val reader = TlsReader(data)
|
||||
|
||||
val version = reader.readUint16()
|
||||
require(version == STATE_VERSION) { "Unsupported state version: $version" }
|
||||
require(version in 1..STATE_VERSION) { "Unsupported state version: $version" }
|
||||
|
||||
val groupContext = GroupContext.decodeTls(reader)
|
||||
val treeBytes = reader.readOpaqueVarInt()
|
||||
@@ -144,6 +170,33 @@ data class MlsGroupState(
|
||||
val interimTranscriptHash = reader.readOpaqueVarInt()
|
||||
val encryptionSecret = reader.readOpaqueVarInt()
|
||||
|
||||
// v2+: per-sender SecretTree ratchet positions. Absent (or an
|
||||
// empty count) for v1 blobs, which restore at generation 0.
|
||||
val senderRatchetStates =
|
||||
if (version >= 2 && reader.hasRemaining) {
|
||||
val count = reader.readUint32().toInt()
|
||||
buildMap {
|
||||
repeat(count) {
|
||||
val leafIndex = reader.readUint32().toInt()
|
||||
val handshakeSecret = reader.readOpaqueVarInt()
|
||||
val handshakeGeneration = reader.readUint32().toInt()
|
||||
val applicationSecret = reader.readOpaqueVarInt()
|
||||
val applicationGeneration = reader.readUint32().toInt()
|
||||
put(
|
||||
leafIndex,
|
||||
SenderRatchetState(
|
||||
handshakeSecret = handshakeSecret,
|
||||
handshakeGeneration = handshakeGeneration,
|
||||
applicationSecret = applicationSecret,
|
||||
applicationGeneration = applicationGeneration,
|
||||
),
|
||||
)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
emptyMap()
|
||||
}
|
||||
|
||||
return MlsGroupState(
|
||||
groupContext = groupContext,
|
||||
treeBytes = treeBytes,
|
||||
@@ -154,6 +207,7 @@ data class MlsGroupState(
|
||||
encryptionPrivateKey = encryptionPrivateKey,
|
||||
interimTranscriptHash = interimTranscriptHash,
|
||||
encryptionSecret = encryptionSecret,
|
||||
senderRatchetStates = senderRatchetStates,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
+26
@@ -389,6 +389,32 @@ class SecretTree(
|
||||
|
||||
return currentSecret
|
||||
}
|
||||
|
||||
/**
|
||||
* Snapshot every sender's current ratchet position so the enclosing
|
||||
* group state can be persisted (RFC 9420 §9).
|
||||
*
|
||||
* Without this, a restore rebuilds the tree at generation 0 for every
|
||||
* sender, and the LOCAL member then re-emits generation 0 within the
|
||||
* same epoch on its next send — reusing the AEAD key+nonce (a
|
||||
* confidentiality break) and getting rejected by strict receivers
|
||||
* (openmls / MDK / Whitenoise) that forbid generation reuse.
|
||||
*
|
||||
* Only the live ratchet position (secret + generation) per sender is
|
||||
* captured. The replay-detection and skipped-key caches are runtime-only
|
||||
* and deliberately excluded — they are safe to drop across a restart.
|
||||
*/
|
||||
fun exportSenderStates(): Map<Int, SenderRatchetState> = senderState.toMap()
|
||||
|
||||
/**
|
||||
* Seed per-sender ratchet positions from an [exportSenderStates]
|
||||
* snapshot. Called by `MlsGroup.restore`. Any sender absent from
|
||||
* [states] simply re-derives from generation 0 on first use, which is
|
||||
* correct for receive-only ratchets.
|
||||
*/
|
||||
fun importSenderStates(states: Map<Int, SenderRatchetState>) {
|
||||
senderState.putAll(states)
|
||||
}
|
||||
}
|
||||
|
||||
data class SenderRatchetState(
|
||||
|
||||
+40
@@ -123,6 +123,46 @@ class MlsGroupManagerTest {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Regression: [MlsGroupManager.encrypt] must persist the advanced ratchet
|
||||
* position, not just commits. A group state persists only at commits was
|
||||
* the second half of the generation-reuse bug: sends between two commits
|
||||
* advanced the SecretTree in memory but never hit the store, so a restart
|
||||
* reloaded the pre-send ratchet and re-emitted an already-used generation.
|
||||
*
|
||||
* Alice and Bob share a group. Alice sends one message (Bob consumes
|
||||
* generation 0), Alice "restarts" from the store WITHOUT any intervening
|
||||
* commit, and her next send must be a fresh generation Bob accepts.
|
||||
*/
|
||||
@Test
|
||||
fun testEncryptPersistsRatchetPositionBetweenCommits() {
|
||||
runBlocking {
|
||||
val aliceStore = InMemoryGroupStateStore()
|
||||
val alice = MlsGroupManager(aliceStore)
|
||||
val aliceGroup = alice.createGroup(groupId, "alice".encodeToByteArray())
|
||||
|
||||
// Bob joins as a low-level MlsGroup — a strict peer that tracks
|
||||
// consumed generations. (The manager's processWelcome requires a
|
||||
// NostrGroupData extension we don't set up here; the low-level
|
||||
// group is enough to observe the ratchet behavior.)
|
||||
val bobBundle = aliceGroup.createKeyPackage("bob".encodeToByteArray(), ByteArray(0))
|
||||
val addResult = alice.addMember(groupId, bobBundle.keyPackage.toTlsBytes())
|
||||
val bob = MlsGroup.processWelcome(addResult.welcomeBytes!!, bobBundle)
|
||||
|
||||
// Alice sends generation 0 (no commit); Bob consumes it.
|
||||
val ct0 = alice.encrypt(groupId, "msg0".encodeToByteArray())
|
||||
assertContentEquals("msg0".encodeToByteArray(), bob.decrypt(ct0).content)
|
||||
|
||||
// Alice restarts from the store — only encrypt() has run since the
|
||||
// last commit, so this proves encrypt persisted the ratchet.
|
||||
val aliceRestarted = MlsGroupManager(aliceStore)
|
||||
aliceRestarted.restoreAll()
|
||||
|
||||
val ct1 = aliceRestarted.encrypt(groupId, "msg1".encodeToByteArray())
|
||||
assertContentEquals("msg1".encodeToByteArray(), bob.decrypt(ct1).content)
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testAddMemberPersistsState() {
|
||||
runBlocking {
|
||||
|
||||
+112
-2
@@ -173,12 +173,12 @@ class MlsGroupStateTest {
|
||||
val state = group.saveState()
|
||||
val bytes = state.encodeTls()
|
||||
|
||||
// First two bytes should be the version (uint16 = 1)
|
||||
// First two bytes should be the version (uint16 = 2)
|
||||
val reader =
|
||||
com.vitorpamplona.quartz.marmot.mls.codec
|
||||
.TlsReader(bytes)
|
||||
val version = reader.readUint16()
|
||||
assertEquals(1, version)
|
||||
assertEquals(2, version)
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -219,4 +219,114 @@ class MlsGroupStateTest {
|
||||
val decrypted = restoredGroup.decrypt(encrypted)
|
||||
assertContentEquals(plaintext, decrypted.content)
|
||||
}
|
||||
|
||||
/**
|
||||
* Regression: a restore must NOT rewind the SecretTree ratchet to
|
||||
* generation 0. A peer that already consumed generation 0 in this epoch
|
||||
* (like openmls / MDK / Whitenoise, which forbid generation reuse) would
|
||||
* otherwise reject the restored sender's next message as a replay.
|
||||
*/
|
||||
@Test
|
||||
fun testRestorePreservesSenderGeneration_peerAcceptsNextMessage() {
|
||||
val alice = MlsGroup.create("alice".encodeToByteArray())
|
||||
val bobBundle =
|
||||
MlsGroup
|
||||
.create("bob".encodeToByteArray())
|
||||
.createKeyPackage("bob".encodeToByteArray(), ByteArray(0))
|
||||
val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle)
|
||||
|
||||
// Alice sends generation 0; Bob consumes it.
|
||||
val ct0 = alice.encrypt("msg0".encodeToByteArray())
|
||||
assertContentEquals("msg0".encodeToByteArray(), bob.decrypt(ct0).content)
|
||||
|
||||
// Alice "restarts": persist then restore.
|
||||
val aliceRestored = MlsGroup.restore(MlsGroupState.decodeTls(alice.saveState().encodeTls()))
|
||||
|
||||
// Alice's next send must be generation 1, which Bob accepts. Before
|
||||
// the fix this re-emitted generation 0 and Bob threw "Generation 0
|
||||
// already consumed".
|
||||
val ct1 = aliceRestored.encrypt("msg1".encodeToByteArray())
|
||||
assertContentEquals("msg1".encodeToByteArray(), bob.decrypt(ct1).content)
|
||||
}
|
||||
|
||||
/**
|
||||
* The ratchet position must survive several sends across a restore, not
|
||||
* just one. Covers the case where the app persists (at a commit) after N
|
||||
* application messages have already advanced the ratchet.
|
||||
*/
|
||||
@Test
|
||||
fun testRestorePreservesSenderGenerationAfterMultipleSends() {
|
||||
val alice = MlsGroup.create("alice".encodeToByteArray())
|
||||
val bobBundle =
|
||||
MlsGroup
|
||||
.create("bob".encodeToByteArray())
|
||||
.createKeyPackage("bob".encodeToByteArray(), ByteArray(0))
|
||||
val bob = MlsGroup.processWelcome(alice.addMember(bobBundle.keyPackage.toTlsBytes()).welcomeBytes!!, bobBundle)
|
||||
|
||||
for (i in 0 until 5) {
|
||||
val ct = alice.encrypt("m$i".encodeToByteArray())
|
||||
assertContentEquals("m$i".encodeToByteArray(), bob.decrypt(ct).content)
|
||||
}
|
||||
|
||||
val aliceRestored = MlsGroup.restore(MlsGroupState.decodeTls(alice.saveState().encodeTls()))
|
||||
|
||||
// Continues at generation 5 — Bob (who consumed 0..4) accepts it.
|
||||
val ct = aliceRestored.encrypt("m5".encodeToByteArray())
|
||||
assertContentEquals("m5".encodeToByteArray(), bob.decrypt(ct).content)
|
||||
}
|
||||
|
||||
/**
|
||||
* Backward compatibility: a STATE_VERSION 1 blob (no persisted ratchet
|
||||
* positions) must still decode, yielding an empty ratchet map and the
|
||||
* legacy generation-0 restore behavior.
|
||||
*/
|
||||
@Test
|
||||
fun testDecodeLegacyV1StateBlob() {
|
||||
val group = MlsGroup.create("alice".encodeToByteArray())
|
||||
group.encrypt("advance the ratchet".encodeToByteArray())
|
||||
val state = group.saveState()
|
||||
|
||||
val v1Bytes = encodeAsV1(state)
|
||||
val decoded = MlsGroupState.decodeTls(v1Bytes)
|
||||
|
||||
assertTrue(decoded.senderRatchetStates.isEmpty(), "v1 blob has no ratchet positions")
|
||||
|
||||
// Restores and can still encrypt/decrypt (legacy behavior).
|
||||
val restored = MlsGroup.restore(decoded)
|
||||
val ct = restored.encrypt("post-restore".encodeToByteArray())
|
||||
assertContentEquals("post-restore".encodeToByteArray(), restored.decrypt(ct).content)
|
||||
}
|
||||
|
||||
/**
|
||||
* Re-encode a state in the original STATE_VERSION 1 layout: identical to
|
||||
* v2 but with the version tag set to 1 and no trailing ratchet section.
|
||||
*/
|
||||
private fun encodeAsV1(state: MlsGroupState): ByteArray {
|
||||
val writer =
|
||||
com.vitorpamplona.quartz.marmot.mls.codec
|
||||
.TlsWriter()
|
||||
writer.putUint16(1)
|
||||
state.groupContext.encodeTls(writer)
|
||||
writer.putOpaqueVarInt(state.treeBytes)
|
||||
writer.putUint32(state.myLeafIndex.toLong())
|
||||
val es = state.epochSecrets
|
||||
writer.putOpaqueVarInt(es.joinerSecret)
|
||||
writer.putOpaqueVarInt(es.welcomeSecret)
|
||||
writer.putOpaqueVarInt(es.epochSecret)
|
||||
writer.putOpaqueVarInt(es.senderDataSecret)
|
||||
writer.putOpaqueVarInt(es.encryptionSecret)
|
||||
writer.putOpaqueVarInt(es.exporterSecret)
|
||||
writer.putOpaqueVarInt(es.epochAuthenticator)
|
||||
writer.putOpaqueVarInt(es.externalSecret)
|
||||
writer.putOpaqueVarInt(es.confirmationKey)
|
||||
writer.putOpaqueVarInt(es.membershipKey)
|
||||
writer.putOpaqueVarInt(es.resumptionPsk)
|
||||
writer.putOpaqueVarInt(es.initSecret)
|
||||
writer.putOpaqueVarInt(state.initSecret)
|
||||
writer.putOpaqueVarInt(state.signingPrivateKey)
|
||||
writer.putOpaqueVarInt(state.encryptionPrivateKey)
|
||||
writer.putOpaqueVarInt(state.interimTranscriptHash)
|
||||
writer.putOpaqueVarInt(state.encryptionSecret)
|
||||
return writer.toByteArray()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user