mirror of
https://github.com/vitorpamplona/amethyst.git
synced 2026-10-06 03:38:23 +00:00
perf(marmot): stop rewriting the whole message log on every send
Recording one Marmot message read the entire conversation back through the Android KeyStore, appended a line, and pushed all of it through again. A room a few thousand messages long was moving hundreds of KB through a hardware-backed cipher to write a couple of hundred bytes - on the send path, and growing with the history. That is the shape of "it used to be fast". The log is now a sequence of independently encrypted segments, so an append encrypts and writes only the new entry. Loose segments fold back into one every 200 appends, which keeps a read from costing one cipher round-trip per message ever sent. Entries are held in memory, so an append no longer reads the log back at all - the same cache answers the duplicate check that used to require decrypting everything. The format lives in EncryptedAppendLog, in commons rather than inside the Android store, because it holds the only copy of a decrypted Marmot message: the ratchet moved past the ciphertext long before anyone reads it back, so a bug here loses history outright, and it needs tests that an Android-only class cannot have. The tests cover the cases involving files this code did not write - one in the original whole-blob format, one whose last append was cut short by process death, one with a segment that will not decrypt. Old files are read as before and upgraded by the next append. The epochs, expiries, retention and snapshot logs share the codec and get the same treatment. Alongside it, KeyStoreEncryption stops fetching the key handle from the keystore daemon on every single operation, which put an IPC in front of every group-state write and every appended message. The handle is cached and dropped if an operation ever fails, so a key the system retires is re-read rather than failing forever. Its one shared Cipher is now one per thread. A Cipher carries the state of the operation in progress, so two coroutines encrypting on different Dispatchers.IO threads through the same instance could corrupt each other's output - latent, but reachable now that more of the send path runs off the caller's thread. Not done here, and worth stating: the per-send MLS group-state write remains, and it is the floor this cannot go below without splitting the ratchet position out of the state blob. Whether any of this is the 4-5 seconds still needs a measurement on a real device - in particular whether AMETHYST_AES_KEY ended up StrongBox-backed, where AES throughput is orders of magnitude lower than the TEE. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_019R58HADiyhsipTye538fWs
This commit is contained in:
+37
-99
@@ -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
|
||||
* <rootDir>/mls_groups/<nostrGroupId>/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<String> =
|
||||
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<String, Long> =
|
||||
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<String, Long> =
|
||||
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<String>,
|
||||
) = 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<Long, Long> =
|
||||
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<String> = readAllFrom(messagesFile(nostrGroupId))
|
||||
|
||||
private fun readAllFrom(file: File): List<String> {
|
||||
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<String>(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<String> = log.readAll(file)
|
||||
|
||||
private fun writeAll(
|
||||
nostrGroupId: String,
|
||||
@@ -359,51 +338,10 @@ class AndroidMarmotMessageStore(
|
||||
private fun writeAllTo(
|
||||
file: File,
|
||||
messages: List<String>,
|
||||
) {
|
||||
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]+$")
|
||||
}
|
||||
}
|
||||
|
||||
+25
-2
@@ -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
|
||||
}
|
||||
|
||||
+40
-12
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+272
@@ -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<String>,
|
||||
val seen: MutableSet<String>,
|
||||
var segments: Int,
|
||||
)
|
||||
|
||||
private val logs = mutableMapOf<String, LogState>()
|
||||
|
||||
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<String> = 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<String>,
|
||||
) {
|
||||
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<List<String>, 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<String>() 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<String>() to UNSEGMENTED
|
||||
return decodeEntries(plain) to UNSEGMENTED
|
||||
}
|
||||
|
||||
val result = ArrayList<String>()
|
||||
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<String>): 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<String> {
|
||||
if (plain.size < 4) return emptyList()
|
||||
var offset = 0
|
||||
val count = readInt(plain, offset)
|
||||
offset += 4
|
||||
|
||||
val result = ArrayList<String>(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
|
||||
}
|
||||
}
|
||||
+203
@@ -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<String>): 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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user