Consolidate JumpToUnreadState

This commit is contained in:
Jenna Vassar
2026-05-01 09:13:40 -07:00
parent 7148f2f28f
commit 200b98b72b
6 changed files with 94 additions and 66 deletions
@@ -14,14 +14,12 @@ import androidx.compose.runtime.MutableState
import androidx.compose.runtime.collectAsState
import androidx.compose.runtime.derivedStateOf
import androidx.compose.runtime.getValue
import androidx.compose.runtime.mutableIntStateOf
import androidx.compose.runtime.mutableStateOf
import androidx.compose.runtime.produceState
import androidx.compose.runtime.remember
import androidx.compose.runtime.rememberCoroutineScope
import androidx.compose.runtime.saveable.rememberSaveable
import androidx.compose.runtime.setValue
import androidx.compose.runtime.snapshots.Snapshot
import dev.zacsweers.metro.Assisted
import dev.zacsweers.metro.AssistedFactory
import dev.zacsweers.metro.AssistedInject
@@ -32,6 +30,7 @@ import io.element.android.features.messages.impl.crypto.sendfailure.resolve.Reso
import io.element.android.features.messages.impl.timeline.components.MessageShieldData
import io.element.android.features.messages.impl.timeline.factories.TimelineItemsFactory
import io.element.android.features.messages.impl.timeline.factories.TimelineItemsFactoryConfig
import io.element.android.features.messages.impl.timeline.model.JumpToUnreadState
import io.element.android.features.messages.impl.timeline.model.NewEventState
import io.element.android.features.messages.impl.timeline.model.TimelineItem
import io.element.android.features.messages.impl.timeline.model.event.isMessageContent
@@ -282,11 +281,14 @@ class TimelinePresenter(
// read marker advances in place — the SDK swaps the marker virtual item to a new position
// without changing the list length, e.g. when [markRoomAsFullyRead] is sent while at the
// bottom of the room.
val readMarkerIndex = remember { mutableIntStateOf(-1) }
val unreadMessagesCount = remember { mutableIntStateOf(0) }
LaunchedEffect(timelineItems) {
val jumpToUnreadState = remember { mutableStateOf<JumpToUnreadState>(JumpToUnreadState.Disabled) }
LaunchedEffect(timelineItems, displayJumpToUnread) {
if (!displayJumpToUnread) {
jumpToUnreadState.value = JumpToUnreadState.Disabled
return@LaunchedEffect
}
val items = timelineItems
withContext(dispatchers.computation) {
jumpToUnreadState.value = withContext(dispatchers.computation) {
var markerIdx = -1
var unread = 0
for ((i, item) in items.withIndex()) {
@@ -298,12 +300,7 @@ class TimelinePresenter(
unread++
}
}
// Apply both writes atomically so consumers never see a half-updated pair
// (e.g. a non-negative markerIdx with the previous unread count).
Snapshot.withMutableSnapshot {
readMarkerIndex.intValue = markerIdx
unreadMessagesCount.intValue = if (markerIdx < 0) 0 else unread
}
if (markerIdx < 0) JumpToUnreadState.NoMarker else JumpToUnreadState.Loaded(markerIdx, unread)
}
}
@@ -356,9 +353,7 @@ class TimelinePresenter(
resolveVerifiedUserSendFailureState = resolveVerifiedUserSendFailureState,
displayThreadSummaries = displayThreadSummaries,
displayFloatingDateBadge = displayFloatingDateBadge,
displayJumpToUnread = displayJumpToUnread,
readMarkerIndex = readMarkerIndex.intValue,
unreadMessagesCount = unreadMessagesCount.intValue,
jumpToUnreadState = jumpToUnreadState.value,
eventSink = ::handleEvent,
)
}
@@ -11,6 +11,7 @@ package io.element.android.features.messages.impl.timeline
import androidx.compose.runtime.Immutable
import io.element.android.features.messages.impl.crypto.sendfailure.resolve.ResolveVerifiedUserSendFailureState
import io.element.android.features.messages.impl.timeline.components.MessageShieldData
import io.element.android.features.messages.impl.timeline.model.JumpToUnreadState
import io.element.android.features.messages.impl.timeline.model.NewEventState
import io.element.android.features.messages.impl.timeline.model.TimelineItem
import io.element.android.features.messages.impl.typing.TypingNotificationState
@@ -35,9 +36,7 @@ data class TimelineState(
val resolveVerifiedUserSendFailureState: ResolveVerifiedUserSendFailureState,
val displayThreadSummaries: Boolean,
val displayFloatingDateBadge: Boolean,
val displayJumpToUnread: Boolean,
val readMarkerIndex: Int,
val unreadMessagesCount: Int,
val jumpToUnreadState: JumpToUnreadState,
val eventSink: (TimelineEvent) -> Unit,
) {
private val lastTimelineEvent = timelineItems.firstOrNull { it is TimelineItem.Event } as? TimelineItem.Event
@@ -12,6 +12,7 @@ import io.element.android.features.messages.impl.crypto.sendfailure.resolve.Reso
import io.element.android.features.messages.impl.crypto.sendfailure.resolve.aResolveVerifiedUserSendFailureState
import io.element.android.features.messages.impl.timeline.components.MessageShieldData
import io.element.android.features.messages.impl.timeline.components.receipt.aReadReceiptData
import io.element.android.features.messages.impl.timeline.model.JumpToUnreadState
import io.element.android.features.messages.impl.timeline.model.NewEventState
import io.element.android.features.messages.impl.timeline.model.ReadReceiptData
import io.element.android.features.messages.impl.timeline.model.TimelineItem
@@ -57,9 +58,7 @@ fun aTimelineState(
resolveVerifiedUserSendFailureState: ResolveVerifiedUserSendFailureState = aResolveVerifiedUserSendFailureState(),
displayThreadSummaries: Boolean = false,
displayFloatingDateBadge: Boolean = false,
displayJumpToUnread: Boolean = true,
readMarkerIndex: Int = -1,
unreadMessagesCount: Int = 0,
jumpToUnreadState: JumpToUnreadState = JumpToUnreadState.NoMarker,
newEventState: NewEventState = NewEventState.None,
eventSink: (TimelineEvent) -> Unit = {},
): TimelineState {
@@ -81,9 +80,7 @@ fun aTimelineState(
resolveVerifiedUserSendFailureState = resolveVerifiedUserSendFailureState,
displayThreadSummaries = displayThreadSummaries,
displayFloatingDateBadge = displayFloatingDateBadge,
displayJumpToUnread = displayJumpToUnread,
readMarkerIndex = readMarkerIndex,
unreadMessagesCount = unreadMessagesCount,
jumpToUnreadState = jumpToUnreadState,
eventSink = eventSink,
)
}
@@ -65,6 +65,7 @@ import io.element.android.features.messages.impl.timeline.components.toText
import io.element.android.features.messages.impl.timeline.di.LocalTimelineItemPresenterFactories
import io.element.android.features.messages.impl.timeline.di.aFakeTimelineItemPresenterFactories
import io.element.android.features.messages.impl.timeline.focus.FocusRequestStateView
import io.element.android.features.messages.impl.timeline.model.JumpToUnreadState
import io.element.android.features.messages.impl.timeline.model.NewEventState
import io.element.android.features.messages.impl.timeline.model.TimelineItem
import io.element.android.features.messages.impl.timeline.model.event.TimelineItemEventContent
@@ -220,9 +221,7 @@ fun TimelineView(
newEventState = state.newEventState,
isLive = state.isLive,
focusRequestState = state.focusRequestState,
readMarkerIndex = state.readMarkerIndex,
unreadMessagesCount = state.unreadMessagesCount,
displayJumpToUnread = state.displayJumpToUnread,
jumpToUnreadState = state.jumpToUnreadState,
topInset = floatingDateTopOffset,
onScrollFinishAt = ::onScrollFinishAt,
onJumpToLive = ::onJumpToLive,
@@ -307,9 +306,7 @@ private fun BoxScope.TimelineScrollHelper(
forceJumpToBottomVisibility: Boolean,
forceJumpToReadMarkerVisibility: Boolean,
focusRequestState: FocusRequestState,
readMarkerIndex: Int,
unreadMessagesCount: Int,
displayJumpToUnread: Boolean,
jumpToUnreadState: JumpToUnreadState,
topInset: Dp,
onScrollFinishAt: (Int) -> Unit,
onJumpToLive: () -> Unit,
@@ -326,10 +323,10 @@ private fun BoxScope.TimelineScrollHelper(
derivedStateOf {
when {
forceJumpToReadMarkerVisibility -> true
!displayJumpToUnread || readMarkerIndex < 0 -> false
jumpToUnreadState !is JumpToUnreadState.Loaded -> false
else -> {
val lastVisibleIndex = lazyListState.layoutInfo.visibleItemsInfo.lastOrNull()?.index ?: return@derivedStateOf false
readMarkerIndex > lastVisibleIndex
jumpToUnreadState.markerIndex > lastVisibleIndex
}
}
}
@@ -362,9 +359,9 @@ private fun BoxScope.TimelineScrollHelper(
}
fun jumpToReadMarker() {
if (readMarkerIndex < 0) return
val loaded = jumpToUnreadState as? JumpToUnreadState.Loaded ?: return
coroutineScope.launch {
lazyListState.animateScrollToItemCenter(readMarkerIndex)
lazyListState.animateScrollToItemCenter(loaded.markerIndex)
}
}
@@ -407,14 +404,15 @@ private fun BoxScope.TimelineScrollHelper(
contentDescription = stringResource(id = CommonStrings.a11y_jump_to_bottom),
isVisible = isJumpToBottomVisible,
// Hide the badge entirely when the feature is off, regardless of the count value.
count = if (displayJumpToUnread) newEventState.messageCount else 0,
count = if (jumpToUnreadState is JumpToUnreadState.Disabled) 0 else newEventState.messageCount,
modifier = Modifier
.align(Alignment.BottomEnd)
.padding(end = 24.dp, bottom = 12.dp),
onClick = ::jumpToBottom,
)
val jumpToUnreadDescription = if (unreadMessagesCount > 0) {
pluralStringResource(CommonPlurals.a11y_jump_to_unread_messages_count, unreadMessagesCount, unreadMessagesCount)
val unreadCount = (jumpToUnreadState as? JumpToUnreadState.Loaded)?.unreadCount ?: 0
val jumpToUnreadDescription = if (unreadCount > 0) {
pluralStringResource(CommonPlurals.a11y_jump_to_unread_messages_count, unreadCount, unreadCount)
} else {
stringResource(id = CommonStrings.a11y_jump_to_unread_messages)
}
@@ -422,7 +420,7 @@ private fun BoxScope.TimelineScrollHelper(
icon = CompoundIcons.ChevronUp(),
contentDescription = jumpToUnreadDescription,
isVisible = isJumpToUnreadVisible,
count = unreadMessagesCount,
count = unreadCount,
// Top padding includes [topInset] so the FAB sits below any pinned-events banner.
modifier = Modifier
.align(Alignment.TopEnd)
@@ -563,8 +561,7 @@ private fun TimelineViewWithReadMarker(
// Index points past the loaded items, mirroring the real-world state the FAB
// represents: the user has scrolled past the read marker, so it's no longer in
// view. The actual scroll target doesn't matter for a static preview.
readMarkerIndex = timelineItems.size,
unreadMessagesCount = unreadMessagesCount,
jumpToUnreadState = JumpToUnreadState.Loaded(markerIndex = timelineItems.size, unreadCount = unreadMessagesCount),
newEventState = if (newMessagesCount > 0) NewEventState.FromOther(newMessagesCount) else NewEventState.None,
),
timelineProtectionState = aTimelineProtectionState(),
@@ -0,0 +1,32 @@
/*
* Copyright (c) 2025 Element Creations Ltd.
* Copyright 2023-2025 New Vector Ltd.
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial.
* Please see LICENSE files in the repository root for full details.
*/
package io.element.android.features.messages.impl.timeline.model
import androidx.compose.runtime.Immutable
/**
* Drives the jump-to-unread FAB and the count badge on the scroll-to-bottom FAB.
*
* The two affordances share state because they're both gated on the same feature flag and both
* use counts derived from the same timeline scan.
*/
@Immutable
sealed interface JumpToUnreadState {
/** Feature flag is off — neither the FAB nor the new-message badge is shown. */
data object Disabled : JumpToUnreadState
/** Feature flag is on, but no read marker is present in the current timeline window. */
data object NoMarker : JumpToUnreadState
/**
* Feature flag is on and the read marker is loaded at [markerIndex]. The FAB shows when the
* marker is above the viewport; the badge displays [unreadCount] when greater than zero.
*/
data class Loaded(val markerIndex: Int, val unreadCount: Int) : JumpToUnreadState
}
@@ -16,6 +16,7 @@ import io.element.android.features.messages.impl.fixtures.aMessageEvent
import io.element.android.features.messages.impl.fixtures.aTimelineItemsFactoryCreator
import io.element.android.features.messages.impl.timeline.components.MessageShieldData
import io.element.android.features.messages.impl.timeline.components.aCriticalShield
import io.element.android.features.messages.impl.timeline.model.JumpToUnreadState
import io.element.android.features.messages.impl.timeline.model.NewEventState
import io.element.android.features.messages.impl.timeline.model.TimelineItem
import io.element.android.features.messages.impl.typing.aTypingNotificationState
@@ -27,6 +28,7 @@ import io.element.android.features.poll.api.actions.SendPollResponseAction
import io.element.android.features.poll.test.actions.FakeEndPollAction
import io.element.android.features.poll.test.actions.FakeSendPollResponseAction
import io.element.android.features.roomcall.api.aStandByCallState
import io.element.android.libraries.featureflag.api.FeatureFlags
import io.element.android.libraries.featureflag.test.FakeFeatureFlagService
import io.element.android.libraries.matrix.api.core.EventId
import io.element.android.libraries.matrix.api.core.RoomId
@@ -372,14 +374,15 @@ class TimelinePresenterTest {
}
@Test
fun `present - unreadMessagesCount counts message-content items between newest and read marker, excluding state events`() = runTest {
fun `present - jumpToUnreadState reports loaded marker and unread count, excluding state events`() = runTest {
val timelineItems = MutableStateFlow(emptyList<MatrixTimelineItem>())
val timeline = FakeTimeline(timelineItems = timelineItems)
val presenter = createTimelinePresenter(timeline)
val presenter = createTimelinePresenter(
timeline = timeline,
featureFlagService = FakeFeatureFlagService(initialState = mapOf(FeatureFlags.JumpToUnread.key to true)),
)
presenter.test {
val initialState = awaitFirstItem()
assertThat(initialState.readMarkerIndex).isEqualTo(-1)
assertThat(initialState.unreadMessagesCount).isEqualTo(0)
awaitFirstItem()
// SDK delivers items oldest-first; the factory reverses so output index 0 is the newest.
// After processing: [msg-newest, membership, msg-2, read-marker, msg-old]
timelineItems.emit(
@@ -391,20 +394,22 @@ class TimelinePresenterTest {
MatrixTimelineItem.Event(UniqueId("msg-newest"), anEventTimelineItem(content = aMessageContent())),
)
)
consumeItemsUntilPredicate { it.readMarkerIndex >= 0 }.last().also { state ->
consumeItemsUntilPredicate { it.jumpToUnreadState is JumpToUnreadState.Loaded }.last().also { state ->
// 2 message items above the marker; the membership state event is skipped.
assertThat(state.readMarkerIndex).isEqualTo(3)
assertThat(state.unreadMessagesCount).isEqualTo(2)
assertThat(state.jumpToUnreadState).isEqualTo(JumpToUnreadState.Loaded(markerIndex = 3, unreadCount = 2))
}
cancelAndIgnoreRemainingEvents()
}
}
@Test
fun `present - unreadMessagesCount excludes own messages and PAGINATION-origin events`() = runTest {
fun `present - jumpToUnreadState count excludes own messages and PAGINATION-origin events`() = runTest {
val timelineItems = MutableStateFlow(emptyList<MatrixTimelineItem>())
val timeline = FakeTimeline(timelineItems = timelineItems)
val presenter = createTimelinePresenter(timeline)
val presenter = createTimelinePresenter(
timeline = timeline,
featureFlagService = FakeFeatureFlagService(initialState = mapOf(FeatureFlags.JumpToUnread.key to true)),
)
presenter.test {
awaitFirstItem()
// After processing (factory reverses): [other-newest, own-msg, paginated, read-marker, msg-old]
@@ -420,33 +425,34 @@ class TimelinePresenterTest {
MatrixTimelineItem.Event(UniqueId("other-newest"), anEventTimelineItem(content = aMessageContent())),
)
)
consumeItemsUntilPredicate { it.readMarkerIndex >= 0 }.last().also { state ->
assertThat(state.readMarkerIndex).isEqualTo(3)
consumeItemsUntilPredicate { it.jumpToUnreadState is JumpToUnreadState.Loaded }.last().also { state ->
// Only `other-newest` counts: own-msg and paginated are filtered out.
assertThat(state.unreadMessagesCount).isEqualTo(1)
assertThat(state.jumpToUnreadState).isEqualTo(JumpToUnreadState.Loaded(markerIndex = 3, unreadCount = 1))
}
cancelAndIgnoreRemainingEvents()
}
}
@Test
fun `present - readMarkerIndex is -1 when no read marker present`() = runTest {
fun `present - jumpToUnreadState is NoMarker when no read marker present`() = runTest {
val timelineItems = MutableStateFlow(emptyList<MatrixTimelineItem>())
val timeline = FakeTimeline(timelineItems = timelineItems)
val presenter = createTimelinePresenter(timeline)
val presenter = createTimelinePresenter(
timeline = timeline,
featureFlagService = FakeFeatureFlagService(initialState = mapOf(FeatureFlags.JumpToUnread.key to true)),
)
presenter.test {
val initialState = awaitFirstItem()
assertThat(initialState.readMarkerIndex).isEqualTo(-1)
assertThat(initialState.unreadMessagesCount).isEqualTo(0)
awaitFirstItem()
timelineItems.emit(
listOf(
MatrixTimelineItem.Event(UniqueId("1"), anEventTimelineItem(content = aMessageContent())),
MatrixTimelineItem.Event(UniqueId("2"), anEventTimelineItem(content = aMessageContent())),
)
)
consumeItemsUntilPredicate { it.timelineItems.size == 2 }.last().also { state ->
assertThat(state.readMarkerIndex).isEqualTo(-1)
assertThat(state.unreadMessagesCount).isEqualTo(0)
consumeItemsUntilPredicate {
it.timelineItems.size == 2 && it.jumpToUnreadState == JumpToUnreadState.NoMarker
}.last().also { state ->
assertThat(state.jumpToUnreadState).isEqualTo(JumpToUnreadState.NoMarker)
}
cancelAndIgnoreRemainingEvents()
}
@@ -614,18 +620,20 @@ class TimelinePresenterTest {
}
@Test
fun `present - readMarkerIndex is 0 when the read marker is the only item`() = runTest {
fun `present - jumpToUnreadState markerIndex is 0 when the read marker is the only item`() = runTest {
val timelineItems = MutableStateFlow(emptyList<MatrixTimelineItem>())
val timeline = FakeTimeline(timelineItems = timelineItems)
val presenter = createTimelinePresenter(timeline)
val presenter = createTimelinePresenter(
timeline = timeline,
featureFlagService = FakeFeatureFlagService(initialState = mapOf(FeatureFlags.JumpToUnread.key to true)),
)
presenter.test {
awaitFirstItem()
timelineItems.emit(
listOf(MatrixTimelineItem.Virtual(UniqueId("read-marker"), VirtualTimelineItem.ReadMarker))
)
consumeItemsUntilPredicate { it.readMarkerIndex >= 0 }.last().also { state ->
assertThat(state.readMarkerIndex).isEqualTo(0)
assertThat(state.unreadMessagesCount).isEqualTo(0)
consumeItemsUntilPredicate { it.jumpToUnreadState is JumpToUnreadState.Loaded }.last().also { state ->
assertThat(state.jumpToUnreadState).isEqualTo(JumpToUnreadState.Loaded(markerIndex = 0, unreadCount = 0))
}
cancelAndIgnoreRemainingEvents()
}