diff --git a/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicImpl.kt b/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicImpl.kt index 5a10b3dd1e7..36632cde573 100644 --- a/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicImpl.kt +++ b/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicImpl.kt @@ -16,6 +16,7 @@ package io.getstream.chat.android.state.plugin.logic.channel.internal +import androidx.annotation.VisibleForTesting import io.getstream.chat.android.client.ChatClient import io.getstream.chat.android.client.api.models.Pagination import io.getstream.chat.android.client.api.models.QueryChannelRequest @@ -28,6 +29,7 @@ import io.getstream.chat.android.client.extensions.internal.NEVER import io.getstream.chat.android.client.extensions.internal.applyPagination import io.getstream.chat.android.client.persistance.repository.RepositoryFacade import io.getstream.chat.android.client.query.pagination.AnyChannelPaginationRequest +import io.getstream.chat.android.client.utils.message.isLocalOnly import io.getstream.chat.android.models.Channel import io.getstream.chat.android.models.Message import io.getstream.chat.android.state.model.querychannels.pagination.internal.QueryChannelPaginationRequest @@ -382,17 +384,19 @@ internal class ChannelLogicImpl( * * @param direction [Pagination] instance which shows direction of pagination. */ - private fun getLoadMoreBaseMessage(direction: Pagination): Message? { - val messages = mutableState.sortedMessages.value.takeUnless(Collection::isEmpty) ?: return null + @VisibleForTesting + internal fun getLoadMoreBaseMessage(direction: Pagination): Message? { + // The server resolves the anchor by id, so a message it does not know about cannot be one + val messages = mutableState.sortedMessages.value return when (direction) { Pagination.GREATER_THAN_OR_EQUAL, Pagination.GREATER_THAN, - -> messages.last() + -> messages.lastOrNull { !it.isLocalOnly() } Pagination.LESS_THAN, Pagination.LESS_THAN_OR_EQUAL, Pagination.AROUND_ID, - -> messages.first() + -> messages.firstOrNull { !it.isLocalOnly() } } } } diff --git a/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogic.kt b/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogic.kt index acedeb43e11..a79f2efc6a6 100644 --- a/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogic.kt +++ b/stream-chat-android-state/src/main/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogic.kt @@ -28,6 +28,7 @@ import io.getstream.chat.android.client.events.UserStopWatchingEvent import io.getstream.chat.android.client.extensions.internal.NEVER import io.getstream.chat.android.client.setup.state.ClientState import io.getstream.chat.android.client.utils.message.isDeleted +import io.getstream.chat.android.client.utils.message.isLocalOnly import io.getstream.chat.android.client.utils.message.isPinExpired import io.getstream.chat.android.client.utils.message.isPinned import io.getstream.chat.android.client.utils.message.isReply @@ -361,7 +362,7 @@ internal class ChannelStateLogic( } messages.filter { it.isReply() }.forEach(::addQuotedMessage) when (shouldRefreshMessages) { - true -> mutableState.setMessages(messages) + true -> mutableState.setMessages(messages + localOnlyMessagesMissingFrom(messages)) else -> { val oldMessages = mutableState.messageList.value.associateBy(Message::id) @@ -375,6 +376,16 @@ internal class ChannelStateLogic( messages.forEach { it.storePoll() } } + /** + * Local only messages are never part of a server response, so a refresh would drop them from the + * list until the channel is reloaded from the database. + */ + private fun localOnlyMessagesMissingFrom(messages: List): List { + val serverMessageIds = messages.mapTo(mutableSetOf(), Message::id) + return mutableState.messageList.value + .filter { it.id !in serverMessageIds && it.isLocalOnly() } + } + /** * Upsert pinned messages in the channel. * diff --git a/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicTest.kt b/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicTest.kt index c54d7821a97..cb1d05c8b52 100644 --- a/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicTest.kt +++ b/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelLogicTest.kt @@ -16,17 +16,24 @@ package io.getstream.chat.android.state.plugin.logic.channel.internal +import io.getstream.chat.android.client.api.models.Pagination import io.getstream.chat.android.client.test.randomMemberAddedEvent import io.getstream.chat.android.client.test.randomMemberRemovedEvent import io.getstream.chat.android.client.test.randomUserMessagesDeletedEvent +import io.getstream.chat.android.models.Message +import io.getstream.chat.android.models.MessageType +import io.getstream.chat.android.models.SyncStatus import io.getstream.chat.android.randomBoolean import io.getstream.chat.android.randomCID import io.getstream.chat.android.randomDate import io.getstream.chat.android.randomMember +import io.getstream.chat.android.randomMessage import io.getstream.chat.android.randomString import io.getstream.chat.android.randomUser import io.getstream.chat.android.state.plugin.state.channel.internal.ChannelMutableState +import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.test.TestScope +import org.amshove.kluent.`should be equal to` import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.mockito.kotlin.doReturn @@ -38,6 +45,7 @@ import org.mockito.kotlin.whenever internal class ChannelLogicTest { private val currentUserId = randomString() + private val sortedMessages = MutableStateFlow>(emptyList()) private lateinit var channelStateLogic: ChannelStateLogic private lateinit var sut: ChannelLogic @@ -47,6 +55,7 @@ internal class ChannelLogicTest { val cid = randomCID() val mutableState = mock() whenever(mutableState.cid).doReturn(cid) + whenever(mutableState.sortedMessages).doReturn(sortedMessages) // Channel state logic channelStateLogic = mock() whenever(channelStateLogic.writeChannelState()).doReturn(mutableState) @@ -60,6 +69,43 @@ internal class ChannelLogicTest { ) } + @Test + fun `When paginating back, Then local only messages are not used as the anchor`() { + val localOnly = randomMessage( + syncStatus = SyncStatus.SYNC_NEEDED, + type = MessageType.REGULAR, + createdAt = null, + createdLocallyAt = null, + ) + val oldestServerMessage = randomMessage(syncStatus = SyncStatus.COMPLETED, type = MessageType.REGULAR) + val newestServerMessage = randomMessage(syncStatus = SyncStatus.COMPLETED, type = MessageType.REGULAR) + sortedMessages.value = listOf(localOnly, oldestServerMessage, newestServerMessage) + + val base = (sut as ChannelLogicImpl).getLoadMoreBaseMessage(Pagination.LESS_THAN) + + base?.id `should be equal to` oldestServerMessage.id + } + + @Test + fun `When paginating forward, Then local only messages are not used as the anchor`() { + val newestServerMessage = randomMessage(syncStatus = SyncStatus.COMPLETED, type = MessageType.REGULAR) + val pending = randomMessage(syncStatus = SyncStatus.SYNC_NEEDED, type = MessageType.REGULAR) + sortedMessages.value = listOf(newestServerMessage, pending) + + val base = (sut as ChannelLogicImpl).getLoadMoreBaseMessage(Pagination.GREATER_THAN) + + base?.id `should be equal to` newestServerMessage.id + } + + @Test + fun `When every message is local only, Then there is no anchor`() { + sortedMessages.value = listOf(randomMessage(syncStatus = SyncStatus.SYNC_NEEDED, type = MessageType.REGULAR)) + + val base = (sut as ChannelLogicImpl).getLoadMoreBaseMessage(Pagination.LESS_THAN) + + base `should be equal to` null + } + @Test fun `When handling MemberAddedEvent for current user, Then channel members and membership are updated`() { // Given diff --git a/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogicTest.kt b/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogicTest.kt index 699f1c95716..b5607b7e022 100644 --- a/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogicTest.kt +++ b/stream-chat-android-state/src/test/java/io/getstream/chat/android/state/plugin/logic/channel/internal/ChannelStateLogicTest.kt @@ -30,6 +30,7 @@ import io.getstream.chat.android.models.Config import io.getstream.chat.android.models.Member import io.getstream.chat.android.models.Message import io.getstream.chat.android.models.MessageType +import io.getstream.chat.android.models.SyncStatus import io.getstream.chat.android.models.User import io.getstream.chat.android.models.toChannelData import io.getstream.chat.android.randomCID @@ -54,12 +55,14 @@ import io.getstream.chat.android.test.TestCoroutineExtension import io.getstream.result.Error import kotlinx.coroutines.flow.MutableStateFlow import org.amshove.kluent.`should be equal to` +import org.amshove.kluent.`should contain` import org.amshove.kluent.`should not be equal to` import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.junit.jupiter.api.extension.RegisterExtension import org.mockito.kotlin.any +import org.mockito.kotlin.argumentCaptor import org.mockito.kotlin.doAnswer import org.mockito.kotlin.doReturn import org.mockito.kotlin.eq @@ -308,6 +311,48 @@ internal class ChannelStateLogicTest { verify(mutableState).setMessages(any()) } + @Test + fun `given refresh messages is true, local only messages should be kept`() { + val pendingMessage = randomMessage( + syncStatus = SyncStatus.SYNC_NEEDED, + type = MessageType.REGULAR, + createdAt = null, + createdLocallyAt = randomDate(), + ) + whenever(mutableState.messageList) doReturn MutableStateFlow(listOf(pendingMessage)) + val serverMessages = listOf(randomMessage(syncStatus = SyncStatus.COMPLETED, type = MessageType.REGULAR)) + val channel: Channel = randomChannel(messages = serverMessages) + val request = QueryChannelRequest().apply { shouldRefresh = true }.withMessages(1) + + channelStateLogic.propagateChannelQuery(channel, request) + + val captor = argumentCaptor>() + verify(mutableState).setMessages(captor.capture()) + captor.firstValue.map(Message::id) `should contain` pendingMessage.id + } + + @Test + fun `given refresh messages is true, the server copy wins over a stale local only one`() { + val id = randomString() + val staleLocalCopy = randomMessage( + id = id, + syncStatus = SyncStatus.SYNC_NEEDED, + type = MessageType.REGULAR, + createdAt = null, + createdLocallyAt = randomDate(), + ) + whenever(mutableState.messageList) doReturn MutableStateFlow(listOf(staleLocalCopy)) + val serverCopy = randomMessage(id = id, syncStatus = SyncStatus.COMPLETED, type = MessageType.REGULAR) + val channel: Channel = randomChannel(messages = listOf(serverCopy)) + val request = QueryChannelRequest().apply { shouldRefresh = true }.withMessages(1) + + channelStateLogic.propagateChannelQuery(channel, request) + + val captor = argumentCaptor>() + verify(mutableState).setMessages(captor.capture()) + captor.firstValue `should be equal to` listOf(serverCopy) + } + @Test fun `given inside search, if message update comes, should update the message`() { _insideSearch.value = true