Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -117,13 +117,13 @@ public class StatePlugin internal constructor(
private val chatClientConfig: ChatClientConfig,
) : Plugin,
QueryMembersListener by QueryMembersListenerState(logic),
QueryChannelsListener by QueryChannelsListenerState(logic, queryingChannelsFree),
QueryChannelsListener by QueryChannelsListenerState(logic, mutableGlobalState, queryingChannelsFree),
QueryGroupedChannelsListener by QueryGroupedChannelsListenerState(
logic = logic,
globalState = mutableGlobalState,
groupedUnreadChannelsUpdater = groupedUnreadChannelsUpdater,
),
QueryChannelListener by QueryChannelListenerState(logic),
QueryChannelListener by QueryChannelListenerState(logic, mutableGlobalState),
ThreadQueryListener by ThreadQueryListenerState(logic),
ChannelMarkReadListener by ChannelMarkReadListenerState(logic),
EditMessageListener by EditMessageListenerState(logic, clientState),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ package io.getstream.chat.android.client.internal.state.plugin.listener.internal

import io.getstream.chat.android.client.api.models.QueryChannelRequest
import io.getstream.chat.android.client.internal.state.plugin.logic.internal.LogicRegistry
import io.getstream.chat.android.client.internal.state.plugin.state.global.internal.MutableGlobalState
import io.getstream.chat.android.client.plugin.listeners.QueryChannelListener
import io.getstream.chat.android.client.utils.stringify
import io.getstream.chat.android.models.Channel
Expand All @@ -28,8 +29,12 @@ import io.getstream.result.Result
* Implementation of [QueryChannelListener] that handles state updates in the SDK.
*
* @param logic [LogicRegistry]
* @param mutableGlobalState [MutableGlobalState]
*/
internal class QueryChannelListenerState(private val logic: LogicRegistry) : QueryChannelListener {
internal class QueryChannelListenerState(
private val logic: LogicRegistry,
private val mutableGlobalState: MutableGlobalState,
) : QueryChannelListener {

private val logger by taggedLogger("QueryChannelListenerS")

Expand Down Expand Up @@ -77,5 +82,6 @@ internal class QueryChannelListenerState(private val logic: LogicRegistry) : Que
"request: $request, result: ${result.stringify { it.cid }}"
}
logic.channel(channelType, channelId).onQueryChannelResult(request, result)
result.onSuccess { channel -> mutableGlobalState.updateChannelDrafts(listOf(channel)) }
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import io.getstream.chat.android.client.api.models.QueryChannelsRequest
import io.getstream.chat.android.client.api.models.QueryChannelsResult
import io.getstream.chat.android.client.internal.state.model.querychannels.pagination.internal.toOfflinePaginationRequest
import io.getstream.chat.android.client.internal.state.plugin.logic.internal.LogicRegistry
import io.getstream.chat.android.client.internal.state.plugin.state.global.internal.MutableGlobalState
import io.getstream.chat.android.client.plugin.listeners.QueryChannelsListener
import io.getstream.result.Result
import kotlinx.coroutines.flow.MutableStateFlow
Expand All @@ -39,6 +40,7 @@ import kotlinx.coroutines.flow.MutableStateFlow
*/
internal class QueryChannelsListenerState(
private val logic: LogicRegistry,
private val mutableGlobalState: MutableGlobalState,
private val queryingChannelsFree: MutableStateFlow<Boolean>,
) : QueryChannelsListener {

Expand All @@ -65,6 +67,7 @@ internal class QueryChannelsListenerState(
}
val channels = result.map(QueryChannelsResult::channels)
queryChannelsLogic.onQueryChannelsResult(channels, request)
channels.onSuccess(mutableGlobalState::updateChannelDrafts)
queryingChannelsFree.value = true
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,8 @@ internal class QueryGroupedChannelsListenerState(
// would never come down.
finishFirstPageLoads(groups.orEmpty().keys - result.value.groups.keys, groups, completed = true)

result.value.groups.values.forEach { group -> globalState.updateChannelDrafts(group.channels) }

// Route each returned group's channels into the per-group state. The captured config lets
// both ChannelListViewModel.loadMoreGroupedChannels and SyncManager.updateGroupedQueryChannels
// reuse the caller's original parameters on paginated and recovery calls respectively.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ package io.getstream.chat.android.client.internal.state.plugin.state.global.inte

import io.getstream.chat.android.client.api.state.GlobalState
import io.getstream.chat.android.client.internal.state.utils.internal.mapState
import io.getstream.chat.android.models.Channel
import io.getstream.chat.android.models.ChannelMute
import io.getstream.chat.android.models.DraftMessage
import io.getstream.chat.android.models.Location
Expand Down Expand Up @@ -130,6 +131,11 @@ internal class MutableGlobalState(
?.let { it.value += (draftMessage.cid to draftMessage) }
}

/** Adds the drafts the server returned with [channels]. */
fun updateChannelDrafts(channels: List<Channel>) {
channels.forEach { channel -> channel.draftMessage?.let(::updateDraftMessage) }
}

fun removeDraftMessage(draftMessage: DraftMessage) {
draftMessage.parentId?.let { parentId ->
_threadDraftMessages?.let { it.value -= parentId }
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -664,6 +664,7 @@ internal class SyncManager(
logger.v {
"[updateActiveQueryChannels] request completed; foundChannels.size: ${foundChannels.size}"
}
mutableGlobalState.updateChannelDrafts(foundChannels)
updatedCids.addAll(foundChannels.map { it.cid })
logger.v { "[updateActiveQueryChannels] updatedCids.size: ${updatedCids.size}" }
}
Expand Down Expand Up @@ -709,6 +710,7 @@ internal class SyncManager(
?.let(logicRegistry::channel)
?.updateDataForChannel(channel, channel.messages.size)
}
mutableGlobalState.updateChannelDrafts(foundChannels)
repos.storeStateForChannels(foundChannels)
val foundCids = foundChannels.mapTo(mutableSetOf(), Channel::cid)
val stillMissingChannelIds = missingChannelIds.filterNot { it.cid in foundCids }
Expand Down Expand Up @@ -785,6 +787,7 @@ internal class SyncManager(
?.let(logicRegistry::channel)
?.updateDataForChannel(channel, channel.messages.size)
}
mutableGlobalState.updateChannelDrafts(foundChannels)
repos.storeStateForChannels(foundChannels)
foundChannels.mapTo(refreshedCids, Channel::cid)
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,7 @@ import org.junit.jupiter.api.Test
import org.mockito.kotlin.any
import org.mockito.kotlin.anyOrNull
import org.mockito.kotlin.argumentCaptor
import org.mockito.kotlin.atLeastOnce
import org.mockito.kotlin.doReturn
import org.mockito.kotlin.eq
import org.mockito.kotlin.mock
Expand Down Expand Up @@ -295,6 +296,29 @@ internal class SyncManagerTest {
assertEquals(createdAt, _syncState.value?.lastSyncedAt)
}

@Test
fun `performSync adds the drafts of the refreshed watched channels to the global state`() =
runTest(testDispatcher) {
/* Given */
val createdAt = localDate()
val rawCreatedAt = streamDateFormatter.format(createdAt)
val watchedChannel = randomChannel(
type = "messaging",
id = "watched",
draftMessage = randomDraftMessage(parentId = null),
)
givenOversizedSyncPayload(eventCount = 3, createdAt = createdAt, rawCreatedAt = rawCreatedAt)
givenWatchedChannels(cids = setOf(watchedChannel.cid), foundChannels = listOf(watchedChannel))

val syncManager = buildSyncManager(eventReplayMaxCount = 2)

/* When */
syncManager.performSync(cids = listOf("1", "2"))

/* Then */
verify(mutableGlobalState).updateChannelDrafts(listOf(watchedChannel))
}

@Test
fun `performSync replays events when the payload size equals the replay limit`() = runTest(testDispatcher) {
/* Given */
Expand Down Expand Up @@ -1158,6 +1182,33 @@ internal class SyncManagerTest {
verify(chatClient, never()).queryChannelsInternal(any())
}

@Test
fun `reconnect should add the drafts of the recovered channel list to the global state`() =
runTest(testDispatcher) {
val createdAt = localDate()
val rawCreatedAt = streamDateFormatter.format(createdAt)
val channel = randomChannel(draftMessage = randomDraftMessage(parentId = null))
val standardQuery: QueryChannelsLogic = mock {
on(it.groupKey()) doReturn null
on(it.recoveryNeeded()) doReturn MutableStateFlow(true)
onBlocking { it.queryFirstPage() } doReturn Result.Success(listOf(channel))
}

whenever(logicRegistry.getActiveQueryChannelsLogic()) doReturn listOf(standardQuery)
whenever(logicRegistry.getActiveChannelsLogic()) doReturn emptyList()
whenever(stateRegistry.getActiveChannelStates()) doReturn emptyMap()
whenever(clientState.isOnline) doReturn true
whenever(repositoryFacade.selectSyncState(user.id)) doReturn null

val syncManager = buildSyncManager()
syncManager.onEvent(connectedEvent(createdAt, rawCreatedAt))
delay(100)
syncManager.onEvent(connectedEvent(createdAt, rawCreatedAt))
delay(100)

verify(mutableGlobalState, atLeastOnce()).updateChannelDrafts(listOf(channel))
}

@Test
fun `updateActiveQueryChannels should skip grouped queries`() =
runTest(testDispatcher) {
Expand Down Expand Up @@ -1428,6 +1479,36 @@ internal class SyncManagerTest {
assertEquals(setOf(cidA, cidB), filter.values)
}

@Test
fun `reconnect should add the drafts of the refreshed channels to the global state`() =
runTest(testDispatcher) {
val createdAt = localDate()
val rawCreatedAt = streamDateFormatter.format(createdAt)
val channel = randomChannel(type = "messaging", id = "a", draftMessage = randomDraftMessage(parentId = null))
val activeState: ChannelState = mock {
on(it.cid) doReturn channel.cid
on(it.recoveryNeeded) doReturn false
}

whenever(logicRegistry.getActiveQueryChannelsLogic()) doReturn emptyList()
whenever(logicRegistry.getActiveChannelsLogic()) doReturn emptyList()
whenever(logicRegistry.channel(any(), any())) doReturn mock<ChannelLogic>()
whenever(stateRegistry.getActiveChannelStates()) doReturn mapOf(ChannelId.fromCid(channel.cid)!! to activeState)
whenever(chatClient.queryChannelsInternal(any())) doReturn TestCall(
Result.Success(QueryChannelsResult(channels = listOf(channel), predefinedFilter = null)),
)
whenever(clientState.isOnline) doReturn true
whenever(repositoryFacade.selectSyncState(user.id)) doReturn null

val syncManager = buildSyncManager()
syncManager.onEvent(connectedEvent(createdAt, rawCreatedAt))
delay(100)
syncManager.onEvent(connectedEvent(createdAt, rawCreatedAt))
delay(100)

verify(mutableGlobalState).updateChannelDrafts(listOf(channel))
}

@Test
fun `on reconnect with multiple grouped queries should pass per-group limits and shared flags`() =
runTest(testDispatcher) {
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
/*
* Copyright (c) 2014-2026 Stream.io Inc. All rights reserved.
*
* Licensed under the Stream License;
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://github.com/GetStream/stream-chat-android/blob/main/LICENSE
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

package io.getstream.chat.android.client.internal.state.plugin.listener.internal

import io.getstream.chat.android.client.api.models.QueryChannelRequest
import io.getstream.chat.android.client.internal.state.plugin.logic.channel.internal.ChannelLogic
import io.getstream.chat.android.client.internal.state.plugin.logic.internal.LogicRegistry
import io.getstream.chat.android.client.internal.state.plugin.state.global.internal.MutableGlobalState
import io.getstream.chat.android.randomChannel
import io.getstream.chat.android.randomDraftMessage
import io.getstream.chat.android.randomString
import io.getstream.result.Error
import io.getstream.result.Result
import kotlinx.coroutines.test.runTest
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Test
import org.mockito.kotlin.any
import org.mockito.kotlin.doReturn
import org.mockito.kotlin.mock

internal class QueryChannelListenerStateTest {

private val logicRegistry: LogicRegistry = mock {
on { channel(any(), any()) } doReturn mock<ChannelLogic>()
}
private val mutableGlobalState = MutableGlobalState(randomString())
private val listener = QueryChannelListenerState(logicRegistry, mutableGlobalState)

@Test
fun `given the server returns the channel draft, it should reach the global state`() = runTest {
val draftMessage = randomDraftMessage(parentId = null)
val channel = randomChannel(draftMessage = draftMessage)

listener.onQueryChannelResult(Result.Success(channel), channel.type, channel.id, QueryChannelRequest())

assertEquals(mapOf(draftMessage.cid to draftMessage), mutableGlobalState.channelDraftMessages.value)
}

@Test
fun `given the query fails, it should leave the global state drafts untouched`() = runTest {
val draftMessage = randomDraftMessage(parentId = null)
mutableGlobalState.updateDraftMessage(draftMessage)

listener.onQueryChannelResult(
Result.Failure(Error.GenericError(randomString())),
randomString(),
randomString(),
QueryChannelRequest(),
)

assertEquals(mapOf(draftMessage.cid to draftMessage), mutableGlobalState.channelDraftMessages.value)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -21,14 +21,18 @@ import io.getstream.chat.android.client.api.models.QueryChannelsRequest
import io.getstream.chat.android.client.api.models.QueryChannelsResult
import io.getstream.chat.android.client.internal.state.plugin.logic.internal.LogicRegistry
import io.getstream.chat.android.client.internal.state.plugin.logic.querychannels.internal.QueryChannelsLogic
import io.getstream.chat.android.client.internal.state.plugin.state.global.internal.MutableGlobalState
import io.getstream.chat.android.models.Channel
import io.getstream.chat.android.models.Filters
import io.getstream.chat.android.models.querysort.QuerySortByField
import io.getstream.chat.android.randomChannel
import io.getstream.chat.android.randomDraftMessage
import io.getstream.chat.android.randomString
import io.getstream.result.Error
import io.getstream.result.Result
import kotlinx.coroutines.flow.MutableStateFlow
import kotlinx.coroutines.test.runTest
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.BeforeEach
import org.junit.jupiter.api.Test
Expand All @@ -44,6 +48,7 @@ internal class QueryChannelsListenerStateTest {
private lateinit var queryChannelsLogic: QueryChannelsLogic
private lateinit var logicRegistry: LogicRegistry
private lateinit var queryingChannelsFree: MutableStateFlow<Boolean>
private lateinit var mutableGlobalState: MutableGlobalState
private lateinit var listener: QueryChannelsListenerState

@BeforeEach
Expand All @@ -53,7 +58,8 @@ internal class QueryChannelsListenerStateTest {
on { queryChannels(any<QueryChannelsRequest>()) } doReturn queryChannelsLogic
}
queryingChannelsFree = MutableStateFlow(true)
listener = QueryChannelsListenerState(logicRegistry, queryingChannelsFree)
mutableGlobalState = MutableGlobalState(randomString())
listener = QueryChannelsListenerState(logicRegistry, mutableGlobalState, queryingChannelsFree)
}

@Test
Expand Down Expand Up @@ -99,6 +105,27 @@ internal class QueryChannelsListenerStateTest {
verify(queryChannelsLogic, never()).applyResolvedSpec(any(), any())
}

@Test
fun `onQueryChannelsResult adds the returned channel drafts to the global state`() = runTest {
val draftMessage = randomDraftMessage(parentId = null)
val result = Result.Success(
QueryChannelsResult(
channels = listOf(randomChannel(draftMessage = draftMessage), randomChannel(draftMessage = null)),
predefinedFilter = null,
),
)

val request = QueryChannelsRequest(
filter = Filters.eq("type", "messaging"),
querySort = QuerySortByField.descByName("last_message_at"),
limit = 30,
)

listener.onQueryChannelsResult(result, request)

assertEquals(mapOf(draftMessage.cid to draftMessage), mutableGlobalState.channelDraftMessages.value)
}

@Test
fun `onQueryChannelsResult does not apply resolved spec on failure`() = runTest {
// Given
Expand Down
Loading
Loading