diff --git a/amethyst/src/main/java/com/vitorpamplona/amethyst/model/marmot/AndroidMarmotMessageStore.kt b/amethyst/src/main/java/com/vitorpamplona/amethyst/model/marmot/AndroidMarmotMessageStore.kt index 7f7874fffe..c179dc5c0d 100644 --- a/amethyst/src/main/java/com/vitorpamplona/amethyst/model/marmot/AndroidMarmotMessageStore.kt +++ b/amethyst/src/main/java/com/vitorpamplona/amethyst/model/marmot/AndroidMarmotMessageStore.kt @@ -20,6 +20,7 @@ */ package com.vitorpamplona.amethyst.model.marmot +import com.vitorpamplona.amethyst.commons.marmot.EncryptedAppendLog import com.vitorpamplona.amethyst.model.preferences.KeyStoreEncryption import com.vitorpamplona.quartz.marmot.mls.group.MarmotMessageStore import com.vitorpamplona.quartz.nip01Core.core.Event @@ -38,24 +39,16 @@ import java.io.File * /mls_groups//messages — encrypted message log * ``` * - * The on-disk format (after decryption) is a sequence of length-prefixed - * UTF-8 entries: - * ``` - * uint32 count - * for each entry: - * uint32 length - * byte[length] utf8 - * ``` - * - * The whole blob is rewritten on each append (via atomic rename) — this - * keeps encryption simple (one GCM nonce per write) and is acceptable for - * conversation-scale histories. + * Every file here is an [EncryptedAppendLog], which owns the on-disk format and + * the migration from the original whole-blob one. Recording a message appends a + * small encrypted segment instead of rewriting the conversation, which is what + * keeps the cost of a send flat as the history grows. */ class AndroidMarmotMessageStore( private val rootDir: File, private val encryption: KeyStoreEncryption = KeyStoreEncryption(), ) : MarmotMessageStore { - private val writeMutex = Mutex() + private val logMutex = Mutex() init { Log.d(TAG) { @@ -76,17 +69,16 @@ class AndroidMarmotMessageStore( nostrGroupId: String, innerEventJson: String, ) = withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { try { - val existing = readAll(nostrGroupId).toMutableList() - if (innerEventJson in existing) { + val file = messagesFile(nostrGroupId) + if (log.contains(file, innerEventJson)) { Log.d(TAG) { "appendMessage($nostrGroupId): duplicate entry skipped" } return@withLock } - existing.add(innerEventJson) - writeAll(nostrGroupId, existing) + log.append(file, innerEventJson) Log.d(TAG) { - "appendMessage($nostrGroupId): now ${existing.size} message(s) persisted" + "appendMessage($nostrGroupId): now ${log.readAll(file).size} message(s) persisted" } } catch (e: Exception) { Log.e(TAG, "appendMessage($nostrGroupId) FAILED: ${e.message}", e) @@ -98,7 +90,7 @@ class AndroidMarmotMessageStore( override suspend fun loadMessages(nostrGroupId: String): List = withContext(Dispatchers.IO) { try { - val messages = readAll(nostrGroupId) + val messages = logMutex.withLock { readAll(nostrGroupId) } Log.d(TAG) { "loadMessages($nostrGroupId): loaded ${messages.size} message(s)" } @@ -111,8 +103,9 @@ class AndroidMarmotMessageStore( override suspend fun delete(nostrGroupId: String) { withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { for (file in listOf(messagesFile(nostrGroupId), epochsFile(nostrGroupId), snapshotFile(nostrGroupId), expiriesFile(nostrGroupId), epochRetentionsFile(nostrGroupId))) { + log.forget(file) if (file.exists() && !file.delete()) { Log.w(TAG) { "delete($nostrGroupId): failed to remove ${file.absolutePath}" } } @@ -138,7 +131,7 @@ class AndroidMarmotMessageStore( innerEventId: String, epoch: Long, ) = withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { try { val line = "$innerEventId $epoch" val existing = readAllFrom(epochsFile(nostrGroupId)).toMutableList() @@ -154,7 +147,8 @@ class AndroidMarmotMessageStore( override suspend fun loadEpochs(nostrGroupId: String): Map = withContext(Dispatchers.IO) { try { - readAllFrom(epochsFile(nostrGroupId)) + logMutex + .withLock { readAllFrom(epochsFile(nostrGroupId)) } .mapNotNull { line -> val parts = line.trim().split(' ') if (parts.size != 2) return@mapNotNull null @@ -185,7 +179,7 @@ class AndroidMarmotMessageStore( innerEventId: String, expiresAtSecs: Long, ) = withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { try { val existing = readAllFrom(expiriesFile(nostrGroupId)).toMutableList() if (existing.any { it.substringBefore(' ') == innerEventId }) return@withLock @@ -200,7 +194,8 @@ class AndroidMarmotMessageStore( override suspend fun loadExpiries(nostrGroupId: String): Map = withContext(Dispatchers.IO) { try { - readAllFrom(expiriesFile(nostrGroupId)) + logMutex + .withLock { readAllFrom(expiriesFile(nostrGroupId)) } .mapNotNull { line -> val parts = line.trim().split(' ') if (parts.size != 2) return@mapNotNull null @@ -225,7 +220,7 @@ class AndroidMarmotMessageStore( innerEventIds: Set, ) = withContext(Dispatchers.IO) { if (innerEventIds.isEmpty()) return@withContext - writeMutex.withLock { + logMutex.withLock { try { val kept = readAll(nostrGroupId).filter { json -> @@ -257,7 +252,7 @@ class AndroidMarmotMessageStore( epoch: Long, retentionSecs: Long, ) = withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { try { val existing = readAllFrom(epochRetentionsFile(nostrGroupId)).toMutableList() if (existing.any { it.substringBefore(' ') == epoch.toString() }) return@withLock @@ -272,7 +267,8 @@ class AndroidMarmotMessageStore( override suspend fun loadEpochRetentions(nostrGroupId: String): Map = withContext(Dispatchers.IO) { try { - readAllFrom(epochRetentionsFile(nostrGroupId)) + logMutex + .withLock { readAllFrom(epochRetentionsFile(nostrGroupId)) } .mapNotNull { line -> val parts = line.trim().split(' ') if (parts.size != 2) return@mapNotNull null @@ -302,7 +298,7 @@ class AndroidMarmotMessageStore( nostrGroupId: String, snapshotJson: String, ) = withContext(Dispatchers.IO) { - writeMutex.withLock { + logMutex.withLock { try { writeAllTo(snapshotFile(nostrGroupId), listOf(snapshotJson)) } catch (e: Exception) { @@ -314,42 +310,25 @@ class AndroidMarmotMessageStore( override suspend fun loadGroupSnapshot(nostrGroupId: String): String? = withContext(Dispatchers.IO) { try { - readAllFrom(snapshotFile(nostrGroupId)).firstOrNull() + logMutex.withLock { readAllFrom(snapshotFile(nostrGroupId)) }.firstOrNull() } catch (e: Exception) { Log.e(TAG, "loadGroupSnapshot($nostrGroupId) FAILED: ${e.message}", e) null } } + // The segmented, constant-time-append log every file here is stored as. + // Guarded by [logMutex]: it caches decrypted entries so an append never has + // to read the log back, and that cache assumes a single owner. + private val log = + EncryptedAppendLog( + encrypt = { encryption.encrypt(it) }, + decrypt = { encryption.decrypt(it) }, + ) + private fun readAll(nostrGroupId: String): List = readAllFrom(messagesFile(nostrGroupId)) - private fun readAllFrom(file: File): List { - if (!file.exists()) return emptyList() - val encrypted = file.readBytes() - val plain = encryption.decrypt(encrypted) ?: return emptyList() - if (plain.size < 4) return emptyList() - - var offset = 0 - val count = - ((plain[offset++].toInt() and 0xFF) shl 24) or - ((plain[offset++].toInt() and 0xFF) shl 16) or - ((plain[offset++].toInt() and 0xFF) shl 8) or - (plain[offset++].toInt() and 0xFF) - - val result = ArrayList(count.coerceAtMost(MAX_MESSAGES)) - for (i in 0 until count) { - if (offset + 4 > plain.size) break - val len = - ((plain[offset++].toInt() and 0xFF) shl 24) or - ((plain[offset++].toInt() and 0xFF) shl 16) or - ((plain[offset++].toInt() and 0xFF) shl 8) or - (plain[offset++].toInt() and 0xFF) - if (len < 0 || offset + len > plain.size) break - result.add(plain.copyOfRange(offset, offset + len).decodeToString()) - offset += len - } - return result - } + private fun readAllFrom(file: File): List = log.readAll(file) private fun writeAll( nostrGroupId: String, @@ -359,51 +338,10 @@ class AndroidMarmotMessageStore( private fun writeAllTo( file: File, messages: List, - ) { - file.parentFile?.mkdirs() - - val encodedEntries = messages.map { it.encodeToByteArray() } - val totalSize = 4 + encodedEntries.sumOf { 4 + it.size } - val buffer = ByteArray(totalSize) - var offset = 0 - - val count = encodedEntries.size - buffer[offset++] = (count shr 24).toByte() - buffer[offset++] = (count shr 16).toByte() - buffer[offset++] = (count shr 8).toByte() - buffer[offset++] = count.toByte() - - for (entry in encodedEntries) { - val len = entry.size - buffer[offset++] = (len shr 24).toByte() - buffer[offset++] = (len shr 16).toByte() - buffer[offset++] = (len shr 8).toByte() - buffer[offset++] = len.toByte() - entry.copyInto(buffer, offset) - offset += len - } - - val encrypted = encryption.encrypt(buffer) - atomicWrite(file, encrypted) - } - - private fun atomicWrite( - target: File, - data: ByteArray, - ) { - val tempFile = File(target.parentFile, "${target.name}.tmp") - tempFile.writeBytes(data) - if (!tempFile.renameTo(target)) { - tempFile.copyTo(target, overwrite = true) - if (!tempFile.delete()) { - Log.w(TAG) { "Failed to delete temp file after copy fallback: ${tempFile.absolutePath}" } - } - } - } + ) = log.rewrite(file, messages) companion object { private const val TAG = "AndroidMarmotMessageStore" - private const val MAX_MESSAGES = 1_000_000 private val HEX_PATTERN = Regex("^[0-9a-fA-F]+$") } } diff --git a/amethyst/src/main/java/com/vitorpamplona/amethyst/model/preferences/KeyStoreEncryption.kt b/amethyst/src/main/java/com/vitorpamplona/amethyst/model/preferences/KeyStoreEncryption.kt index 6fcd096a8e..74d9dc7af4 100644 --- a/amethyst/src/main/java/com/vitorpamplona/amethyst/model/preferences/KeyStoreEncryption.kt +++ b/amethyst/src/main/java/com/vitorpamplona/amethyst/model/preferences/KeyStoreEncryption.kt @@ -45,10 +45,27 @@ class KeyStoreEncryption { private const val GCM_TAG_LENGTH_BITS = 128 } - private val cipher = Cipher.getInstance(TRANSFORMATION) + // One Cipher per thread rather than one shared instance. A Cipher carries + // the state of the operation in progress, so two coroutines encrypting on + // different Dispatchers.IO threads through the same object would corrupt + // each other's output. + private val ciphers = ThreadLocal.withInitial { Cipher.getInstance(TRANSFORMATION) } + private val keyStore = KeyStore.getInstance(ANDROID_KEY_STORE).apply { load(null) } - private fun getKey(): SecretKey { + // The key handle never changes for the life of the alias, but fetching it + // is a round trip to the keystore daemon. Every encrypted store here reads + // and writes through this class, so doing that per operation put an IPC in + // front of every Marmot group-state write and every message appended. + @Volatile + private var cachedKey: SecretKey? = null + + private fun getKey(): SecretKey = + cachedKey ?: synchronized(this) { + cachedKey ?: loadOrCreateKey().also { cachedKey = it } + } + + private fun loadOrCreateKey(): SecretKey { val existingKey = keyStore.getEntry(KEY_ALIAS, null) as? KeyStore.SecretKeyEntry return existingKey?.secretKey ?: createKey() } @@ -95,11 +112,15 @@ class KeyStoreEncryption { fun encrypt(bytes: ByteArray): ByteArray { try { // Initializes the cipher in encrypt mode and encrypts data + val cipher = ciphers.get() cipher.init(Cipher.ENCRYPT_MODE, getKey()) val iv = cipher.iv val encrypted = cipher.doFinal(bytes) return iv + encrypted } catch (e: Exception) { + // A key the system has retired (a wipe, a credential reset) keeps + // failing until it is re-read, so the cached handle goes with it. + cachedKey = null Log.e(TAG, "encrypt() failed: ${e.message}", e) throw e } @@ -112,9 +133,11 @@ class KeyStoreEncryption { // IvParameterSpec), so we must pass the 128-bit auth tag length. val iv = bytes.copyOfRange(0, GCM_IV_LENGTH) val data = bytes.copyOfRange(GCM_IV_LENGTH, bytes.size) + val cipher = ciphers.get() cipher.init(Cipher.DECRYPT_MODE, getKey(), GCMParameterSpec(GCM_TAG_LENGTH_BITS, iv)) return cipher.doFinal(data) } catch (e: Exception) { + cachedKey = null Log.e(TAG, "decrypt() failed (input ${bytes.size} bytes): ${e.message}", e) throw e } diff --git a/commons/src/androidMain/kotlin/com/vitorpamplona/amethyst/commons/keystorage/KeyStoreEncryption.kt b/commons/src/androidMain/kotlin/com/vitorpamplona/amethyst/commons/keystorage/KeyStoreEncryption.kt index 81f333fe47..fc0ea04b04 100644 --- a/commons/src/androidMain/kotlin/com/vitorpamplona/amethyst/commons/keystorage/KeyStoreEncryption.kt +++ b/commons/src/androidMain/kotlin/com/vitorpamplona/amethyst/commons/keystorage/KeyStoreEncryption.kt @@ -40,10 +40,24 @@ internal class KeyStoreEncryption { private const val KEY_ALIAS = "AMETHYST_AES_KEY" } - private val cipher = Cipher.getInstance(TRANSFORMATION) + // One Cipher per thread rather than one shared instance: a Cipher holds the + // state of the operation in progress, so two callers on different threads + // through the same object would corrupt each other's output. + private val ciphers = ThreadLocal.withInitial { Cipher.getInstance(TRANSFORMATION) } + private val keyStore = KeyStore.getInstance("AndroidKeyStore").apply { load(null) } - private fun getKey(): SecretKey { + // The handle never changes for the life of the alias, but fetching it is a + // round trip to the keystore daemon — not something to pay per operation. + @Volatile + private var cachedKey: SecretKey? = null + + private fun getKey(): SecretKey = + cachedKey ?: synchronized(this) { + cachedKey ?: loadOrCreateKey().also { cachedKey = it } + } + + private fun loadOrCreateKey(): SecretKey { val existingKey = keyStore.getEntry(KEY_ALIAS, null) as? KeyStore.SecretKeyEntry return existingKey?.secretKey ?: createKey() } @@ -85,18 +99,32 @@ internal class KeyStoreEncryption { private fun createKey(): SecretKey = createKeyStrongBoxIfAvailable() ?: createKeyRegular() fun encrypt(bytes: ByteArray): ByteArray { - // Initializes the cipher in encrypt mode and encrypts data - cipher.init(Cipher.ENCRYPT_MODE, getKey()) - val iv = cipher.iv - val encrypted = cipher.doFinal(bytes) - return iv + encrypted + try { + // Initializes the cipher in encrypt mode and encrypts data + val cipher = ciphers.get() + cipher.init(Cipher.ENCRYPT_MODE, getKey()) + val iv = cipher.iv + val encrypted = cipher.doFinal(bytes) + return iv + encrypted + } catch (e: Exception) { + // A key the system has retired (a wipe, a credential reset) keeps + // failing until it is re-read, so the cached handle goes with it. + cachedKey = null + throw e + } } fun decrypt(bytes: ByteArray): ByteArray { - // Extracts IV and decrypts the data - val iv = bytes.copyOfRange(0, 12) // GCM mode uses 12-byte IV - val data = bytes.copyOfRange(12, bytes.size) - cipher.init(Cipher.DECRYPT_MODE, getKey(), IvParameterSpec(iv)) - return cipher.doFinal(data) + try { + // Extracts IV and decrypts the data + val iv = bytes.copyOfRange(0, 12) // GCM mode uses 12-byte IV + val data = bytes.copyOfRange(12, bytes.size) + val cipher = ciphers.get() + cipher.init(Cipher.DECRYPT_MODE, getKey(), IvParameterSpec(iv)) + return cipher.doFinal(data) + } catch (e: Exception) { + cachedKey = null + throw e + } } } diff --git a/commons/src/jvmAndroid/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLog.kt b/commons/src/jvmAndroid/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLog.kt new file mode 100644 index 0000000000..962157ac1b --- /dev/null +++ b/commons/src/jvmAndroid/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLog.kt @@ -0,0 +1,272 @@ +/* + * 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.amethyst.commons.marmot + +import java.io.File +import java.io.FileOutputStream + +/** + * An encrypted-at-rest log of UTF-8 entries that can be appended to in + * constant time. + * + * The file is a sequence of independently encrypted SEGMENTS: + * ``` + * file := MAGIC segment* + * segment := uint32 encLen, byte[encLen] // whatever [encrypt] produces + * plain := uint32 count, (uint32 len, byte[len])* + * ``` + * + * Segments are the whole point. The format this replaces was a single blob + * covering the entire history, so recording one line meant pushing every line + * ever written back through the cipher and out to disk again — work that grew + * with the log and, for Marmot's message log, was paid on the send path. A + * conversation a few thousand messages long was moving hundreds of KB through + * a hardware-backed cipher to append a couple of hundred bytes. + * + * Appending writes one small segment. Loose segments are folded back into one + * every [compactAfterSegments] appends, which bounds what a read costs: without + * that, an old conversation would need one cipher round-trip per message ever + * sent. + * + * A file written by the older format has no magic prefix and is read as a + * single legacy blob; the next append rewrites it in this format. Nothing else + * migrates it, and nothing needs to — reading handles both. + * + * **Not thread-safe.** Entries are cached in memory so an append never has to + * read the log back, and that cache assumes one owner. Callers hold their own + * lock around every method (see `AndroidMarmotMessageStore`), and one instance + * must own any given file. + * + * @param encrypt must produce a self-describing blob — it carries its own IV / + * nonce, since every segment is encrypted separately. + * @param decrypt returns null for a segment it cannot open; that segment's + * entries are skipped and the rest of the log is still read. + */ +class EncryptedAppendLog( + private val encrypt: (ByteArray) -> ByteArray, + private val decrypt: (ByteArray) -> ByteArray?, + private val compactAfterSegments: Int = COMPACT_AFTER_SEGMENTS, +) { + /** Entries of one file, plus what it takes to append without re-reading it. */ + private class LogState( + val entries: MutableList, + val seen: MutableSet, + var segments: Int, + ) + + private val logs = mutableMapOf() + + private fun stateFor(file: File): LogState = + logs.getOrPut(file.absolutePath) { + val (entries, segments) = decodeFile(file) + LogState(entries.toMutableList(), entries.toMutableSet(), segments) + } + + /** Every entry in [file], oldest first. */ + fun readAll(file: File): List = stateFor(file).entries.toList() + + /** Whether [entry] is already in [file], without reading it back from disk. */ + fun contains( + file: File, + entry: String, + ): Boolean = entry in stateFor(file).seen + + /** Append one entry. Constant time, apart from a periodic compaction. */ + fun append( + file: File, + entry: String, + ) { + val state = stateFor(file) + state.entries.add(entry) + state.seen.add(entry) + + // Either there is no header to append after (a file that does not + // exist yet, or one still in the old format), or too many loose + // segments have piled up to keep reads cheap. A rewrite fixes both, + // and is what lays the header down. + if (state.segments == UNSEGMENTED || state.segments >= compactAfterSegments) { + rewrite(file, state.entries.toList()) + return + } + + file.parentFile?.mkdirs() + val segment = encrypt(encodeEntries(listOf(entry))) + FileOutputStream(file, true).use { out -> + out.write(lengthPrefix(segment.size)) + out.write(segment) + // An append that survives only in the page cache would lose a + // message the UI has already shown as sent. + out.fd.sync() + } + state.segments += 1 + } + + /** Replace the whole log with [entries], as a single segment. */ + fun rewrite( + file: File, + entries: List, + ) { + file.parentFile?.mkdirs() + + val segment = encrypt(encodeEntries(entries)) + val out = ByteArray(MAGIC.size + 4 + segment.size) + MAGIC.copyInto(out, 0) + lengthPrefix(segment.size).copyInto(out, MAGIC.size) + segment.copyInto(out, MAGIC.size + 4) + atomicWrite(file, out) + + val state = logs.getOrPut(file.absolutePath) { LogState(mutableListOf(), mutableSetOf(), 1) } + state.entries.clear() + state.entries.addAll(entries) + state.seen.clear() + state.seen.addAll(entries) + state.segments = 1 + } + + /** Drop the in-memory cache for [file]; call when the file is deleted. */ + fun forget(file: File) { + logs.remove(file.absolutePath) + } + + /** Every entry in [file], and how many segments they came from. */ + private fun decodeFile(file: File): Pair, Int> { + // A file that does not exist yet has no header, so the first append has + // to write one rather than tack a bare segment onto nothing. + if (!file.exists()) return emptyList() to UNSEGMENTED + val bytes = file.readBytes() + + if (!bytes.startsWithMagic()) { + // The older format: the file is one encrypted blob and nothing else. + val plain = decrypt(bytes) ?: return emptyList() to UNSEGMENTED + return decodeEntries(plain) to UNSEGMENTED + } + + val result = ArrayList() + var offset = MAGIC.size + var segments = 0 + while (offset + 4 <= bytes.size) { + val encLen = readInt(bytes, offset) + offset += 4 + // A truncated tail is a half-finished append (process death between + // the write and the sync). Everything before it is intact and is + // what we keep; the torn record is dropped rather than failing the + // whole log. + if (encLen <= 0 || offset + encLen > bytes.size) break + val plain = decrypt(bytes.copyOfRange(offset, offset + encLen)) + offset += encLen + segments += 1 + if (plain != null) result.addAll(decodeEntries(plain)) + } + return result to segments.coerceAtLeast(1) + } + + private fun encodeEntries(entries: List): ByteArray { + val encoded = entries.map { it.encodeToByteArray() } + val buffer = ByteArray(4 + encoded.sumOf { 4 + it.size }) + var offset = 0 + + writeInt(buffer, offset, encoded.size) + offset += 4 + + for (entry in encoded) { + writeInt(buffer, offset, entry.size) + offset += 4 + entry.copyInto(buffer, offset) + offset += entry.size + } + return buffer + } + + private fun decodeEntries(plain: ByteArray): List { + if (plain.size < 4) return emptyList() + var offset = 0 + val count = readInt(plain, offset) + offset += 4 + + val result = ArrayList(count.coerceIn(0, MAX_ENTRIES)) + for (i in 0 until count) { + if (offset + 4 > plain.size) break + val len = readInt(plain, offset) + offset += 4 + if (len < 0 || offset + len > plain.size) break + result.add(plain.copyOfRange(offset, offset + len).decodeToString()) + offset += len + } + return result + } + + /** Write via a temp file and rename, so a crash can't leave a half-written log. */ + private fun atomicWrite( + target: File, + data: ByteArray, + ) { + val tempFile = File(target.parentFile, "${target.name}.tmp") + tempFile.writeBytes(data) + if (!tempFile.renameTo(target)) { + tempFile.copyTo(target, overwrite = true) + tempFile.delete() + } + } + + private fun ByteArray.startsWithMagic(): Boolean { + if (size < MAGIC.size) return false + for (i in MAGIC.indices) if (this[i] != MAGIC[i]) return false + return true + } + + private fun lengthPrefix(value: Int): ByteArray = ByteArray(4).also { writeInt(it, 0, value) } + + private fun writeInt( + target: ByteArray, + offset: Int, + value: Int, + ) { + target[offset] = (value shr 24).toByte() + target[offset + 1] = (value shr 16).toByte() + target[offset + 2] = (value shr 8).toByte() + target[offset + 3] = value.toByte() + } + + private fun readInt( + source: ByteArray, + offset: Int, + ): Int = + ((source[offset].toInt() and 0xFF) shl 24) or + ((source[offset + 1].toInt() and 0xFF) shl 16) or + ((source[offset + 2].toInt() and 0xFF) shl 8) or + (source[offset + 3].toInt() and 0xFF) + + companion object { + /** Marks the segmented format; a file without it predates it. */ + private val MAGIC = "MRMTLOG2".encodeToByteArray() + + /** + * Segment count standing for "this file has no header yet" — either it + * does not exist, or it was written by the original whole-blob format. + * The next append rewrites it rather than appending into nothing. + */ + private const val UNSEGMENTED = -1 + + private const val COMPACT_AFTER_SEGMENTS = 200 + + private const val MAX_ENTRIES = 1_000_000 + } +} diff --git a/commons/src/jvmTest/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLogTest.kt b/commons/src/jvmTest/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLogTest.kt new file mode 100644 index 0000000000..166d81f9fa --- /dev/null +++ b/commons/src/jvmTest/kotlin/com/vitorpamplona/amethyst/commons/marmot/EncryptedAppendLogTest.kt @@ -0,0 +1,203 @@ +/* + * 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.amethyst.commons.marmot + +import java.io.File +import java.nio.file.Files +import kotlin.test.Test +import kotlin.test.assertContentEquals +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertTrue + +/** + * The on-disk format behind Marmot's message log. + * + * This is the only copy of a decrypted Marmot message — the ratchet has moved + * past the ciphertext it came from long before anyone reads it back — so a bug + * here loses conversation history outright. The cases that matter are the ones + * involving a file this code did not write itself: one left by the original + * whole-blob format, and one whose last append never finished. + */ +class EncryptedAppendLogTest { + // A stand-in for AES-GCM: prefixes a "nonce" so each blob is self-describing + // and differs per call, exactly like the real one. + private var nonce = 0 + + private fun cipher(compactAfter: Int = 200) = + EncryptedAppendLog( + encrypt = { plain -> byteArrayOf(NONCE_MARK, (nonce++ % 251).toByte()) + plain.map { (it.toInt() xor 0x5A).toByte() } }, + decrypt = { blob -> + if (blob.size < 2 || blob[0] != NONCE_MARK) { + null + } else { + blob.copyOfRange(2, blob.size).map { (it.toInt() xor 0x5A).toByte() }.toByteArray() + } + }, + compactAfterSegments = compactAfter, + ) + + private fun tempFile(): File = File(Files.createTempDirectory("appendlog").toFile(), "log") + + @Test + fun `entries survive a round trip through a fresh reader`() { + val file = tempFile() + val writer = cipher() + writer.append(file, "first") + writer.append(file, "second") + writer.append(file, "third") + + // A second instance reads only what reached the disk. + assertContentEquals(listOf("first", "second", "third"), cipher().readAll(file)) + } + + @Test + fun `an empty log reads as empty rather than failing`() { + assertContentEquals(emptyList(), cipher().readAll(tempFile())) + } + + @Test + fun `contains answers without reading the file back`() { + val file = tempFile() + val log = cipher() + log.append(file, "hello") + + assertTrue(log.contains(file, "hello")) + assertFalse(log.contains(file, "goodbye")) + } + + @Test + fun `a rewrite replaces the whole log`() { + val file = tempFile() + val log = cipher() + log.append(file, "a") + log.append(file, "b") + log.rewrite(file, listOf("b")) + + assertContentEquals(listOf("b"), cipher().readAll(file)) + assertFalse(cipher().contains(file, "a")) + } + + @Test + fun `compaction keeps every entry and collapses the segments`() { + val file = tempFile() + val log = cipher(compactAfter = 4) + val written = (1..20).map { "entry-$it" } + written.forEach { log.append(file, it) } + + assertContentEquals(written, cipher().readAll(file)) + // Four appends per compaction, so the file can never carry 20 segments' + // worth of framing — it is the bound on read cost that matters here. + assertTrue(file.length() < 20 * SEGMENT_OVERHEAD_CEILING, "log should have been compacted, was ${file.length()} bytes") + } + + @Test + fun `a log written by the original whole-blob format is still readable`() { + val file = tempFile() + file.parentFile.mkdirs() + // Exactly what the old writer produced: one encrypted blob, no magic. + file.writeBytes(legacyBlob(listOf("old-one", "old-two"))) + + assertContentEquals(listOf("old-one", "old-two"), cipher().readAll(file)) + } + + @Test + fun `appending to a legacy log upgrades it without losing anything`() { + val file = tempFile() + file.parentFile.mkdirs() + file.writeBytes(legacyBlob(listOf("old-one", "old-two"))) + + val log = cipher() + log.append(file, "new-one") + + assertContentEquals(listOf("old-one", "old-two", "new-one"), cipher().readAll(file)) + // and the upgraded file is in the new format, so the next append is cheap + assertTrue(file.readBytes().decodeToString().startsWith("MRMTLOG2")) + } + + @Test + fun `a torn final append costs only the torn entry`() { + val file = tempFile() + val log = cipher() + log.append(file, "kept-one") + log.append(file, "kept-two") + + // Simulate process death partway through writing the third segment. + val intact = file.readBytes() + log.append(file, "lost") + val torn = file.readBytes() + file.writeBytes(torn.copyOfRange(0, intact.size + 6)) + + assertContentEquals(listOf("kept-one", "kept-two"), cipher().readAll(file)) + } + + @Test + fun `a segment that cannot be decrypted does not hide the rest`() { + val file = tempFile() + val log = cipher() + log.append(file, "before") + log.append(file, "after") + + // Corrupt the first segment's nonce marker so decrypt returns null for it. + val bytes = file.readBytes() + bytes[MAGIC_LEN + 4] = 0 + file.writeBytes(bytes) + + assertContentEquals(listOf("after"), cipher().readAll(file)) + } + + @Test + fun `entries keep their bytes through the round trip`() { + val file = tempFile() + val log = cipher() + val awkward = listOf("", "emoji 👩‍👧 here", "a\nb\tc", "\"quoted\": {\"json\": 1}") + awkward.forEach { log.append(file, it) } + + assertEquals(awkward, cipher().readAll(file)) + } + + /** The pre-segment on-disk shape: `encrypt(uint32 count, (uint32 len, bytes)*)`. */ + private fun legacyBlob(entries: List): ByteArray { + val encoded = entries.map { it.encodeToByteArray() } + val plain = ByteArray(4 + encoded.sumOf { 4 + it.size }) + var offset = 0 + + fun putInt(value: Int) { + plain[offset++] = (value shr 24).toByte() + plain[offset++] = (value shr 16).toByte() + plain[offset++] = (value shr 8).toByte() + plain[offset++] = value.toByte() + } + putInt(encoded.size) + for (entry in encoded) { + putInt(entry.size) + entry.copyInto(plain, offset) + offset += entry.size + } + return byteArrayOf(NONCE_MARK, 0) + plain.map { (it.toInt() xor 0x5A).toByte() } + } + + companion object { + private const val NONCE_MARK: Byte = 0x7F + private const val MAGIC_LEN = 8 + private const val SEGMENT_OVERHEAD_CEILING = 64 + } +}