mirror of
https://github.com/signalapp/Signal-Android.git
synced 2026-09-20 00:35:47 +01:00
Fix message duplication when editing a message twice in quick succession.
This commit is contained in:
committed by
Greyson Parrelli
parent
05973d9f02
commit
ca9765e69f
@@ -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<Long?>): 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<Long>()
|
||||
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)
|
||||
|
||||
@@ -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 <T> 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})")
|
||||
|
||||
Reference in New Issue
Block a user