From ca9765e69f9be67ce10e2a65e31c92cdf43fcd5b Mon Sep 17 00:00:00 2001 From: Cody Henthorne Date: Mon, 14 Sep 2026 16:28:01 -0400 Subject: [PATCH] Fix message duplication when editing a message twice in quick succession. --- .../securesms/database/MessageTable.kt | 34 ++++- .../database/EditMessageRevisionTest.kt | 143 ++++++++++++++++++ 2 files changed, 169 insertions(+), 8 deletions(-) diff --git a/app/src/main/java/org/thoughtcrime/securesms/database/MessageTable.kt b/app/src/main/java/org/thoughtcrime/securesms/database/MessageTable.kt index 0210555a89..9797ff4cbc 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/database/MessageTable.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/database/MessageTable.kt @@ -2277,11 +2277,30 @@ open class MessageTable(context: Context?, databaseHelper: SignalDatabase) : Dat return readableDatabase.select(ID) .from(TABLE_NAME) - .where("($TABLE_NAME.$ID = ? OR $TABLE_NAME.$ORIGINAL_MESSAGE_ID = ?) AND $TABLE_NAME.$ID < ?", originalMessageId, originalMessageId, messageId) + .where("$ORIGINAL_MESSAGE_ID = ? AND $ID < ?", originalMessageId, messageId) .orderBy("$ID DESC") .limit(1) .run() + .readToSingleLong(originalMessageId) + } + + /** Building an edit off anything but the newest revision leaves the chain with more than one visible message. */ + private fun getLatestRevisionId(messageId: Long): Long { + val chainOriginalId = readableDatabase + .select(ORIGINAL_MESSAGE_ID) + .from(TABLE_NAME) + .where(ID_WHERE, messageId) + .run() .readToSingleLong(0) + + return readableDatabase + .select(ID) + .from(TABLE_NAME) + .where("$ORIGINAL_MESSAGE_ID = ?", if (chainOriginalId > 0) chainOriginalId else messageId) + .orderBy("$ID DESC") + .limit(1) + .run() + .readToSingleLong(messageId) } fun getMessages(messageIds: Collection): MmsReader { @@ -3576,7 +3595,7 @@ open class MessageTable(context: Context?, databaseHelper: SignalDatabase) : Dat var editedMessage: MessageRecord? = null if (message.isMessageEdit) { try { - editedMessage = getMessageRecord(message.messageToEdit) + editedMessage = getMessageRecord(getLatestRevisionId(message.messageToEdit)) if (!MessageConstraintsUtil.isValidEditMessageSend(editedMessage)) { throw MmsException("Message is not valid to edit") } @@ -3727,12 +3746,11 @@ open class MessageTable(context: Context?, databaseHelper: SignalDatabase) : Dat .values(QUOTE_ID to message.sentTimeMillis) .where("$QUOTE_ID = ?", editedMessage.dateSent) .run() - } - if (message.messageToEdit > 0) { + val chainOriginalId = editedMessage.getOriginalOrOwnMessageId().id writableDatabase.update(TABLE_NAME) .values(LATEST_REVISION_ID to messageId) - .where("$ID_WHERE OR $LATEST_REVISION_ID = ?", message.messageToEdit, message.messageToEdit) + .where("$ID != ? AND ($ID = ? OR $ORIGINAL_MESSAGE_ID = ?)", messageId, chainOriginalId, chainOriginalId) .run() val textAttachments = (editedMessage as? MmsMessageRecord)?.slideDeck?.asAttachments()?.filter { it.contentType == MediaUtil.LONG_TEXT }?.mapNotNull { (it as? DatabaseAttachment)?.attachmentId?.id } ?: emptyList() @@ -3740,10 +3758,10 @@ open class MessageTable(context: Context?, databaseHelper: SignalDatabase) : Dat val excludeIds = HashSet() excludeIds += textAttachments excludeIds += linkPreviewAttachments - attachments.duplicateAttachmentsForMessage(messageId, message.messageToEdit, excludeIds) + attachments.duplicateAttachmentsForMessage(messageId, editedMessage.id, excludeIds) - reactions.moveReactionsToNewMessage(messageId, message.messageToEdit) - movePinnedDetailsToNewMessage(newMessageId = messageId, previousId = message.messageToEdit) + reactions.moveReactionsToNewMessage(messageId, editedMessage.id) + movePinnedDetailsToNewMessage(newMessageId = messageId, previousId = editedMessage.id) } val hasCollapsed = maybeCollapseMessage(db = writableDatabase, messageId = messageId, threadId = threadId, dateReceived = dateReceived, messageExtras = message.messageExtras, messageType = type) diff --git a/app/src/test/java/org/thoughtcrime/securesms/database/EditMessageRevisionTest.kt b/app/src/test/java/org/thoughtcrime/securesms/database/EditMessageRevisionTest.kt index ef6a57dc3d..1c566f329c 100644 --- a/app/src/test/java/org/thoughtcrime/securesms/database/EditMessageRevisionTest.kt +++ b/app/src/test/java/org/thoughtcrime/securesms/database/EditMessageRevisionTest.kt @@ -6,10 +6,12 @@ package org.thoughtcrime.securesms.database import android.app.Application +import android.database.Cursor import assertk.assertThat import assertk.assertions.isEqualTo import assertk.assertions.isNotNull import assertk.assertions.isNull +import io.mockk.every import org.junit.Assert.assertEquals import org.junit.Before import org.junit.Rule @@ -18,10 +20,15 @@ import org.junit.runner.RunWith import org.robolectric.RobolectricTestRunner import org.robolectric.annotation.Config import org.signal.core.util.CursorUtil +import org.thoughtcrime.securesms.database.model.MessageId import org.thoughtcrime.securesms.database.model.MmsMessageRecord +import org.thoughtcrime.securesms.database.model.ReactionRecord import org.thoughtcrime.securesms.mms.IncomingMessage +import org.thoughtcrime.securesms.mms.OutgoingMessage +import org.thoughtcrime.securesms.recipients.Recipient import org.thoughtcrime.securesms.recipients.RecipientId import org.thoughtcrime.securesms.testutil.RecipientTestRule +import org.thoughtcrime.securesms.util.RemoteConfig @RunWith(RobolectricTestRunner::class) @Config(manifest = Config.NONE, application = Application::class) @@ -32,11 +39,18 @@ class EditMessageRevisionTest { private lateinit var senderId: RecipientId private var threadId: Long = 0 + private lateinit var contactId: RecipientId + private var contactThreadId: Long = 0 @Before fun setUp() { + every { RemoteConfig.regularDeleteThreshold } returns 86_400L + every { RemoteConfig.adminDeleteThreshold } returns 86_400L + senderId = recipients.createRecipient("Sender Name") threadId = SignalDatabase.threads.getOrCreateThreadIdFor(senderId, false, ThreadTable.DistributionTypes.DEFAULT) + contactId = recipients.createRecipient("Contact Name") + contactThreadId = SignalDatabase.threads.getOrCreateThreadIdFor(Recipient.resolved(contactId)) } @Test @@ -186,6 +200,135 @@ class EditMessageRevisionTest { SignalDatabase.writableDatabase.execSQL("PRAGMA foreign_keys=ON") } + @Test + fun outgoingSequentialEditsChainProperly() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + val edit2Id = insertOutgoingEdit(messageToEdit = edit1Id, sentTimeMillis = 2002) + + assertThat(getLatestRevisionId(originalId)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit1Id)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit2Id)).isNull() + + assertEquals("Exactly one visible revision should remain", 1, countVisibleRevisions(originalId)) + } + + @Test + fun outgoingEditTargetingStaleRevisionDoesNotDuplicate() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + val edit2Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2002) + + assertThat(getLatestRevisionId(originalId)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit1Id)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit2Id)).isNull() + + assertEquals("A stale-target edit must not produce a duplicate visible revision", 1, countVisibleRevisions(originalId)) + } + + @Test + fun outgoingEditTargetingStaleRevisionNumbersFromLatestRevision() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + val edit2Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2002) + + assertEquals("First edit is revision 1", 1, getRevisionNumber(edit1Id)) + assertEquals("Second edit builds on the first, not on the original", 2, getRevisionNumber(edit2Id)) + assertEquals("Every revision points back at the chain root", originalId, getOriginalMessageId(edit2Id)) + } + + @Test + fun outgoingEditTargetingStaleRevisionMovesReactionsToNewRevision() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + SignalDatabase.reactions.addReaction(MessageId(edit1Id), ReactionRecord(emoji = "\uD83D\uDC4D", author = contactId, dateSent = 2001, dateReceived = 2001)) + + val edit2Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2002) + + assertEquals("Reactions must follow the chain onto the visible revision", 1, SignalDatabase.reactions.getReactions(MessageId(edit2Id)).size) + assertEquals("Reactions must not be stranded on a hidden revision", 0, SignalDatabase.reactions.getReactions(MessageId(edit1Id)).size) + } + + @Test + fun outgoingEditRecoversFromChainWithMultipleVisibleRevisions() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + val edit2Id = insertOutgoingEdit(messageToEdit = edit1Id, sentTimeMillis = 2002) + + SignalDatabase.writableDatabase.execSQL( + "UPDATE ${MessageTable.TABLE_NAME} SET ${MessageTable.LATEST_REVISION_ID} = NULL WHERE ${MessageTable.ID} = ?", + arrayOf(edit1Id) + ) + assertEquals("Precondition: chain is damaged", 2, countVisibleRevisions(originalId)) + + val edit3Id = insertOutgoingEdit(messageToEdit = edit1Id, sentTimeMillis = 2003) + + assertEquals("A damaged chain must collapse back to a single visible revision", 1, countVisibleRevisions(originalId)) + assertThat(getLatestRevisionId(originalId)).isNotNull().isEqualTo(edit3Id) + assertThat(getLatestRevisionId(edit1Id)).isNotNull().isEqualTo(edit3Id) + assertThat(getLatestRevisionId(edit2Id)).isNotNull().isEqualTo(edit3Id) + assertThat(getLatestRevisionId(edit3Id)).isNull() + } + + @Test + fun outgoingEditCollapsesVisibleChainRoot() { + val originalId = recipients.insertOutgoingMessage(contactId, body = "original", sentTimeMillis = 2000) + val edit1Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2001) + + SignalDatabase.writableDatabase.execSQL( + "UPDATE ${MessageTable.TABLE_NAME} SET ${MessageTable.LATEST_REVISION_ID} = NULL WHERE ${MessageTable.ID} = ?", + arrayOf(originalId) + ) + assertEquals("Precondition: chain root is visible", 2, countVisibleRevisions(originalId)) + + val edit2Id = insertOutgoingEdit(messageToEdit = originalId, sentTimeMillis = 2002) + + assertEquals("A visible chain root must be collapsed too", 1, countVisibleRevisions(originalId)) + assertThat(getLatestRevisionId(originalId)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit1Id)).isNotNull().isEqualTo(edit2Id) + assertThat(getLatestRevisionId(edit2Id)).isNull() + } + + private fun insertOutgoingEdit(messageToEdit: Long, sentTimeMillis: Long): Long { + val message = OutgoingMessage( + recipient = Recipient.resolved(contactId), + body = "edited at $sentTimeMillis", + timestamp = sentTimeMillis, + isSecure = true, + messageToEdit = messageToEdit + ) + return recipients.insertOutgoingMessage(message, contactThreadId) + } + + private fun countVisibleRevisions(originalId: Long): Int { + return SignalDatabase.writableDatabase + .query( + "SELECT COUNT(*) FROM ${MessageTable.TABLE_NAME} WHERE ${MessageTable.LATEST_REVISION_ID} IS NULL AND (${MessageTable.ID} = ? OR ${MessageTable.ORIGINAL_MESSAGE_ID} = ?)", + arrayOf(originalId, originalId) + ) + .use { cursor -> + cursor.moveToFirst() + cursor.getInt(0) + } + } + + private fun getRevisionNumber(messageId: Long): Int { + return readColumn(messageId, MessageTable.REVISION_NUMBER) { cursor, idx -> cursor.getInt(idx) } + } + + private fun getOriginalMessageId(messageId: Long): Long { + return readColumn(messageId, MessageTable.ORIGINAL_MESSAGE_ID) { cursor, idx -> cursor.getLong(idx) } + } + + private fun readColumn(messageId: Long, column: String, read: (Cursor, Int) -> T): T { + return SignalDatabase.writableDatabase + .query(MessageTable.TABLE_NAME, arrayOf(column), "${MessageTable.ID} = ?", arrayOf(messageId.toString()), null, null, null) + .use { cursor -> + require(cursor.moveToFirst()) { "No message with id $messageId" } + read(cursor, cursor.getColumnIndexOrThrow(column)) + } + } + private fun countDanglingLatestRevisionIds(): Int { return SignalDatabase.writableDatabase .query("SELECT COUNT(*) FROM ${MessageTable.TABLE_NAME} WHERE ${MessageTable.LATEST_REVISION_ID} IS NOT NULL AND ${MessageTable.LATEST_REVISION_ID} NOT IN (SELECT ${MessageTable.ID} FROM ${MessageTable.TABLE_NAME})")