From f141bd536bb8ffe0395f497bfa53c41bbd90ab0d Mon Sep 17 00:00:00 2001 From: Cody Henthorne Date: Mon, 17 Aug 2026 16:26:45 -0400 Subject: [PATCH] Fix memory pressure and record thrashing when exporting chat items. --- .../v2/exporters/ChatItemArchiveExporter.kt | 90 +++++-- .../exporters/ChatItemArchiveExporterTest.kt | 240 ++++++++++++++++++ 2 files changed, 311 insertions(+), 19 deletions(-) create mode 100644 app/src/test/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporterTest.kt diff --git a/app/src/main/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporter.kt b/app/src/main/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporter.kt index d5915c1c45..b7f48cb226 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporter.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporter.kt @@ -6,6 +6,7 @@ package org.thoughtcrime.securesms.backup.v2.exporters import android.database.Cursor +import androidx.annotation.VisibleForTesting import okio.ByteString.Companion.toByteString import org.json.JSONArray import org.json.JSONException @@ -142,12 +143,19 @@ class ChatItemArchiveExporter( private val backupStartTime: Long, private val batchSize: Int, private val exportState: ExportState, - private val cursorGenerator: (Long, Int) -> Cursor + private val cursorGenerator: (Long, Int) -> Cursor, + private val maxBufferMemorySize: Int = MAX_BUFFER_MEMORY_SIZE ) : Iterator, Closeable { companion object { val EXPIRATION_CUTOFF = 1.days private val MAX_BUFFER_MEMORY_SIZE = 15.mb + + /** Never ask the database for fewer rows than this, no matter how large the individual records are. */ + private const val MIN_ROW_LIMIT = 100 + + /** How many extra rows to ask for beyond what we expect to consume, to account for records getting smaller. */ + private const val ROW_LIMIT_HEADROOM = 1.5 } /** Timer for more macro-level events, like fetching extra data vs transforming the data. */ @@ -169,7 +177,19 @@ class ChatItemArchiveExporter( private var lastSeenReceivedTime = 0L - private var records: LinkedHashMap = readNextMessageRecordBatch(emptySet()) + /** + * The ids of every record we've already exported that shares [lastSeenReceivedTime]. + */ + private val lastSeenReceivedTimeIds: MutableSet = hashSetOf() + + /** + * The number of rows we ask the database for when reading the next batch. Starts with max and then adjusts + * up and down to account for changes in message sizes. + */ + private var rowLimit = batchSize + + @VisibleForTesting + internal var records: LinkedHashMap = readNextMessageRecordBatch() override fun hasNext(): Boolean { return buffer.isNotEmpty() || records.isNotEmpty() @@ -480,10 +500,9 @@ class ChatItemArchiveExporter( } eventTimer.emit("transform") - val recordIds = HashSet(records.keys) records.clear() - records = readNextMessageRecordBatch(recordIds) + records = readNextMessageRecordBatch() eventTimer.emit("messages") return if (buffer.isNotEmpty()) { @@ -499,22 +518,55 @@ class ChatItemArchiveExporter( Log.d(TAG, "[ChatItemArchiveExporterExtraData][batchSize = $batchSize] ${extraDataTimer.stop().summary}") } - private fun readNextMessageRecordBatch(pastIds: Set): LinkedHashMap { - return cursorGenerator(lastSeenReceivedTime, batchSize).use { cursor -> - val records: LinkedHashMap = LinkedHashMap(batchSize) - var estimatedRecordsMemorySize = 0 - while (cursor.moveToNext() && estimatedRecordsMemorySize < MAX_BUFFER_MEMORY_SIZE) { - cursor.toBackupMessageRecord(pastIds, backupStartTime)?.let { record -> - records[record.id] = record - lastSeenReceivedTime = record.dateReceived - estimatedRecordsMemorySize += record.estimatedSizeInBytes + @VisibleForTesting + internal fun readNextMessageRecordBatch(): LinkedHashMap { + var limit = rowLimit + + while (true) { + val batch: LinkedHashMap = LinkedHashMap(limit.coerceAtMost(batchSize)) + var estimatedBatchMemorySize = 0 + var rowsRead = 0 + + cursorGenerator(lastSeenReceivedTime, limit).use { cursor -> + while (cursor.moveToNext()) { + rowsRead++ + + val record = cursor.toBackupMessageRecord(lastSeenReceivedTimeIds, backupStartTime) ?: continue + + if (record.dateReceived != lastSeenReceivedTime) { + lastSeenReceivedTime = record.dateReceived + lastSeenReceivedTimeIds.clear() + } + lastSeenReceivedTimeIds += record.id + + batch[record.id] = record + estimatedBatchMemorySize += record.estimatedSizeInBytes + + if (estimatedBatchMemorySize >= maxBufferMemorySize) { + break + } } } - if (estimatedRecordsMemorySize > MAX_BUFFER_MEMORY_SIZE) { - Log.d(TAG, "[readNextMessageRecordBatch] recordsSize = ${records.size} recordsMemSize: ${estimatedRecordsMemorySize.bytes.toUnitString(spaced = false)}") + if (batch.isEmpty() && rowsRead >= limit) { + limit = lastSeenReceivedTimeIds.size + MIN_ROW_LIMIT + Log.w(TAG, "[readNextMessageRecordBatch] All $rowsRead rows read were already exported. Retrying with a limit of $limit.") + continue } - records + + if (batch.isNotEmpty()) { + val previousRowLimit = rowLimit + val averageRecordSize = max(1, estimatedBatchMemorySize / batch.size) + val sizeBasedLimit = ((maxBufferMemorySize / averageRecordSize) * ROW_LIMIT_HEADROOM).toInt() + + rowLimit = sizeBasedLimit.coerceIn(MIN_ROW_LIMIT, max(MIN_ROW_LIMIT, batchSize)) + lastSeenReceivedTimeIds.size + + if (rowLimit != previousRowLimit) { + Log.d(TAG, "[readNextMessageRecordBatch] recordsSize = ${batch.size}, recordsMemSize = ${estimatedBatchMemorySize.bytes.toUnitString(spaced = false)}, avgRecordSize = ${averageRecordSize.bytes.toUnitString(spaced = false)}. Adjusting rowLimit $previousRowLimit -> $rowLimit") + } + } + + return batch } } @@ -1827,9 +1879,9 @@ private fun RecipientId.hasAciOrE164(exportState: ExportState): Boolean { return exportState.recipientIdToAci[this.toLong()] != null || exportState.recipientIdToE164[this.toLong()] != null } -private fun Cursor.toBackupMessageRecord(pastIds: Set, backupStartTime: Long): BackupMessageRecord? { +private fun Cursor.toBackupMessageRecord(skipIds: Set, backupStartTime: Long): BackupMessageRecord? { val id = this.requireLong(MessageTable.ID) - if (pastIds.contains(id)) { + if (skipIds.contains(id)) { return null } @@ -1879,7 +1931,7 @@ private fun Cursor.toBackupMessageRecord(pastIds: Set, backupStartTime: Lo ) } -private class BackupMessageRecord( +internal class BackupMessageRecord( val id: Long, val dateSent: Long, val dateReceived: Long, diff --git a/app/src/test/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporterTest.kt b/app/src/test/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporterTest.kt new file mode 100644 index 0000000000..7ffa511a55 --- /dev/null +++ b/app/src/test/java/org/thoughtcrime/securesms/backup/v2/exporters/ChatItemArchiveExporterTest.kt @@ -0,0 +1,240 @@ +/* + * Copyright 2026 Signal Messenger, LLC + * SPDX-License-Identifier: AGPL-3.0-only + */ + +package org.thoughtcrime.securesms.backup.v2.exporters + +import android.app.Application +import android.database.Cursor +import android.database.MatrixCursor +import assertk.assertThat +import assertk.assertions.hasSize +import assertk.assertions.isEmpty +import assertk.assertions.isEqualTo +import assertk.assertions.isGreaterThan +import assertk.assertions.isLessThan +import io.mockk.mockk +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner +import org.robolectric.annotation.Config +import org.thoughtcrime.securesms.backup.v2.ExportState +import org.thoughtcrime.securesms.database.MessageTable +import org.thoughtcrime.securesms.database.SignalDatabase +import org.thoughtcrime.securesms.recipients.RecipientId + +@RunWith(RobolectricTestRunner::class) +@Config(application = Application::class) +class ChatItemArchiveExporterTest { + + @Test + fun `all records exported exactly once when memory limit ends every batch`() { + val rows = (1L..1000L).map { FakeRow(id = it, dateReceived = it, bodySize = LARGE_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 1000).drain() + + assertThat(exported).isEqualTo(rows.map { it.id }) + assertThat(table.requestedCounts.size).isGreaterThan(1) + } + + @Test + fun `all records exported exactly once when row limit ends every batch`() { + val rows = (1L..1000L).map { FakeRow(id = it, dateReceived = it, bodySize = SMALL_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 100).drain() + + assertThat(exported).isEqualTo(rows.map { it.id }) + } + + @Test(timeout = 30_000) + fun `all records exported exactly once when every record shares a dateReceived`() { + val rows = (1L..500L).map { FakeRow(id = it, dateReceived = 1000L, bodySize = LARGE_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 1000).drain() + + assertThat(exported).isEqualTo(rows.map { it.id }) + } + + @Test(timeout = 30_000) + fun `all records exported exactly once when large groups of records share a dateReceived`() { + val rows = (1L..1000L).map { FakeRow(id = it, dateReceived = it / 300L, bodySize = LARGE_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 1000).drain() + + assertThat(exported).isEqualTo(rows.map { it.id }) + } + + @Test + fun `asks for fewer rows after memory pressure and more once records shrink`() { + val large = (1L..600L).map { FakeRow(id = it, dateReceived = it, bodySize = LARGE_BODY) } + val small = (601L..3000L).map { FakeRow(id = it, dateReceived = it, bodySize = SMALL_BODY) } + val table = FakeMessageTable(large + small) + + val batches = table.exporter(batchSize = BATCH_SIZE).drainBatches() + + assertThat(table.requestedCounts.first()).isEqualTo(BATCH_SIZE) + assertThat(table.requestedCounts.min()).isLessThan(BATCH_SIZE) + assertThat(table.requestedCounts.last()).isGreaterThan(table.requestedCounts.min()) + assertThat(table.requestedCounts.min()).isGreaterThan(MIN_ROW_LIMIT - 1) + + assertThat(batches.first().size).isLessThan(batches.last().size) + } + + @Test + fun `no batch exceeds the memory limit by more than one record`() { + val rows = (1L..1000L).map { FakeRow(id = it, dateReceived = it, bodySize = LARGE_BODY) } + val exporter = FakeMessageTable(rows).exporter(batchSize = 1000) + + var batch = exporter.records + while (batch.isNotEmpty()) { + val batchSize = batch.values.sumOf { it.estimatedSizeInBytes } + val largestRecord = batch.values.maxOf { it.estimatedSizeInBytes } + assertThat(batchSize).isLessThan(MAX_MEMORY + largestRecord) + batch = exporter.readNextMessageRecordBatch() + } + } + + @Test + fun `empty table produces no records`() { + val exporter = FakeMessageTable(emptyList()).exporter(batchSize = 100) + + assertThat(exporter.records).isEmpty() + assertThat(exporter.hasNext()).isEqualTo(false) + } + + @Test(timeout = 30_000) + fun `all records exported exactly once when the requested batch size is degenerate`() { + val rows = (1L..500L).map { FakeRow(id = it, dateReceived = it, bodySize = SMALL_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 0).drain() + + assertThat(exported).isEqualTo(rows.map { it.id }) + } + + @Test + fun `records within a batch are ordered by dateReceived`() { + val rows = (1L..300L).map { FakeRow(id = 301L - it, dateReceived = it, bodySize = SMALL_BODY) } + val table = FakeMessageTable(rows) + + val exported = table.exporter(batchSize = 100).drain() + + assertThat(exported).hasSize(300) + assertThat(exported).isEqualTo(rows.sortedBy { it.dateReceived }.map { it.id }) + } + + private fun ChatItemArchiveExporter.drain(): List { + return drainBatches().flatten() + } + + private fun ChatItemArchiveExporter.drainBatches(): List> { + val batches = mutableListOf>() + var batch = records + var reads = 0 + + while (batch.isNotEmpty()) { + batches += batch.keys.toList() + check(++reads < MAX_READS) { "Read $reads batches without exhausting the data. Export is not making progress." } + batch = readNextMessageRecordBatch() + } + + return batches + } + + private data class FakeRow(val id: Long, val dateReceived: Long, val bodySize: Int) + + /** + * Stands in for the real export query, which returns rows with `date_received >= lastSeenReceivedTime` ordered by + * `date_received` ascending, capped at the requested count. + */ + private class FakeMessageTable(private val rows: List) { + val requestedCounts = mutableListOf() + + fun exporter(batchSize: Int): ChatItemArchiveExporter { + return ChatItemArchiveExporter( + db = mockk(relaxed = true), + selfRecipientId = RecipientId.from(1), + noteToSelfThreadId = 1, + backupStartTime = 0, + batchSize = batchSize, + exportState = mockk(relaxed = true), + cursorGenerator = ::query, + maxBufferMemorySize = MAX_MEMORY + ) + } + + private fun query(lastSeenReceivedTime: Long, count: Int): Cursor { + requestedCounts += count + + val cursor = MatrixCursor(COLUMNS) + rows + .filter { it.dateReceived >= lastSeenReceivedTime } + .sortedBy { it.dateReceived } + .take(count) + .forEach { row -> + cursor.newRow() + .add(MessageTable.ID, row.id) + .add(MessageTable.DATE_RECEIVED, row.dateReceived) + .add(MessageTable.DATE_SENT, row.dateReceived) + .add(MessageTable.BODY, "b".repeat(row.bodySize)) + } + + return cursor + } + } + + companion object { + private const val MAX_MEMORY = 200_000 + private const val BATCH_SIZE = 1000 + private const val LARGE_BODY = 1000 + private const val SMALL_BODY = 10 + private const val MAX_READS = 10_000 + + /** Mirrors ChatItemArchiveExporter.MIN_ROW_LIMIT, which is private. */ + private const val MIN_ROW_LIMIT = 100 + + private val COLUMNS = arrayOf( + MessageTable.ID, + MessageTable.DATE_SENT, + MessageTable.DATE_RECEIVED, + MessageTable.DATE_SERVER, + MessageTable.TYPE, + MessageTable.THREAD_ID, + MessageTable.BODY, + MessageTable.MESSAGE_RANGES, + MessageTable.FROM_RECIPIENT_ID, + MessageTable.TO_RECIPIENT_ID, + MessageTable.EXPIRES_IN, + MessageTable.EXPIRE_STARTED, + MessageTable.UNIDENTIFIED, + MessageTable.LINK_PREVIEWS, + MessageTable.SHARED_CONTACTS, + MessageTable.QUOTE_ID, + MessageTable.QUOTE_AUTHOR, + MessageTable.QUOTE_BODY, + MessageTable.QUOTE_MISSING, + MessageTable.QUOTE_BODY_RANGES, + MessageTable.QUOTE_TYPE, + MessageTable.ORIGINAL_MESSAGE_ID, + MessageTable.LATEST_REVISION_ID, + MessageTable.HAS_DELIVERY_RECEIPT, + MessageTable.VIEWED_COLUMN, + MessageTable.HAS_READ_RECEIPT, + MessageTable.READ, + MessageTable.RECEIPT_TIMESTAMP, + MessageTable.NETWORK_FAILURES, + MessageTable.MISMATCHED_IDENTITIES, + MessageTable.MESSAGE_EXTRAS, + MessageTable.VIEW_ONCE, + MessageTable.PARENT_STORY_ID, + MessageTable.PINNED_AT, + MessageTable.PINNED_UNTIL, + MessageTable.DELETED_BY + ) + } +}