feat: adjust event propagation to avoid android primitives (+perf)

KMP sure is coming
This commit is contained in:
infi 2026-07-29 06:03:01 +02:00
parent 771cef8e80
commit 12e71c1bc6
4 changed files with 239 additions and 111 deletions

View File

@ -1,7 +1,5 @@
package chat.stoat.api package chat.stoat.api
import android.os.Handler
import android.os.Looper
import android.util.Log import android.util.Log
import androidx.compose.runtime.mutableStateMapOf import androidx.compose.runtime.mutableStateMapOf
import chat.stoat.BuildConfig import chat.stoat.BuildConfig
@ -38,22 +36,29 @@ import io.ktor.client.plugins.websocket.WebSockets
import io.ktor.client.request.header import io.ktor.client.request.header
import io.ktor.serialization.kotlinx.json.json import io.ktor.serialization.kotlinx.json.json
import io.sentry.Sentry import io.sentry.Sentry
import kotlinx.coroutines.CancellationException
import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.DelicateCoroutinesApi import kotlinx.coroutines.DelicateCoroutinesApi
import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.ExperimentalCoroutinesApi
import kotlinx.coroutines.Job import kotlinx.coroutines.Job
import kotlinx.coroutines.cancelAndJoin
import kotlinx.coroutines.delay
import kotlinx.coroutines.flow.MutableSharedFlow import kotlinx.coroutines.flow.MutableSharedFlow
import kotlinx.coroutines.isActive
import kotlinx.coroutines.launch import kotlinx.coroutines.launch
import kotlinx.coroutines.newSingleThreadContext import kotlinx.coroutines.newSingleThreadContext
import kotlinx.coroutines.runBlocking
import kotlinx.coroutines.withContext import kotlinx.coroutines.withContext
import kotlinx.serialization.ExperimentalSerializationApi import kotlinx.serialization.ExperimentalSerializationApi
import kotlinx.serialization.SerialName import kotlinx.serialization.SerialName
import kotlinx.serialization.Serializable import kotlinx.serialization.Serializable
import kotlinx.serialization.cbor.Cbor import kotlinx.serialization.cbor.Cbor
import kotlinx.serialization.json.Json import kotlinx.serialization.json.Json
import logcat.LogPriority
import logcat.asLog
import logcat.logcat
import java.net.SocketException import java.net.SocketException
import kotlin.time.Duration.Companion.seconds
import chat.stoat.core.model.schemas.Channel as ChannelSchema import chat.stoat.core.model.schemas.Channel as ChannelSchema
fun String.api(): String { fun String.api(): String {
@ -132,10 +137,13 @@ val StoatHttp = HttpClient(OkHttp) {
} }
} }
val mainHandler = Handler(Looper.getMainLooper())
object StoatAPI { object StoatAPI {
const val TOKEN_HEADER_NAME = "x-session-token" const val TOKEN_HEADER_NAME = "x-session-token"
private const val WS_EVENT_BUFFER_CAPACITY =
128 // arbitrary -- should be adjusted if too much gets dropped...
private val INITIAL_RECONNECT_DELAY = 1.seconds
private val MAX_RECONNECT_DELAY = 30.seconds
private val PING_INTERVAL = 30.seconds // Same interval as the web clients (/revolt.js)
val userCache = mutableStateMapOf<String, User>() val userCache = mutableStateMapOf<String, User>()
val serverCache = mutableStateMapOf<String, Server>() val serverCache = mutableStateMapOf<String, Server>()
@ -159,10 +167,11 @@ object StoatAPI {
val realtimeContext = newSingleThreadContext("RealtimeContext") val realtimeContext = newSingleThreadContext("RealtimeContext")
val wsFrameChannel = MutableSharedFlow<Any>( val wsFrameChannel = MutableSharedFlow<Any>(
replay = 0, replay = 0,
extraBufferCapacity = Int.MAX_VALUE, extraBufferCapacity = WS_EVENT_BUFFER_CAPACITY,
) )
private var socketCoroutine: Job? = null private var socketCoroutine: Job? = null
private var pingCoroutine: Job? = null
private var openForLocalHydration = true private var openForLocalHydration = true
@ -183,28 +192,36 @@ object StoatAPI {
@OptIn(ExperimentalCoroutinesApi::class) @OptIn(ExperimentalCoroutinesApi::class)
suspend fun connectWS() { suspend fun connectWS() {
socketCoroutine?.cancelAndJoin()
RealtimeSocket.updateDisconnectionState(DisconnectionState.Reconnecting)
val token = sessionToken
socketCoroutine = CoroutineScope(Dispatchers.IO).launch { socketCoroutine = CoroutineScope(Dispatchers.IO).launch {
try { var reconnectDelay = INITIAL_RECONNECT_DELAY
withContext(realtimeContext) { while (isActive && sessionToken == token) {
try {
RealtimeSocket.connect(sessionToken)
} catch (e: SocketException) {
Log.d("RevoltAPI", "Socket closed, probably no big deal /// " + e.message)
RealtimeSocket.updateDisconnectionState(DisconnectionState.Disconnected)
} catch (e: Exception) {
Log.e("RevoltAPI", "WebSocket error", e)
RealtimeSocket.updateDisconnectionState(DisconnectionState.Disconnected)
}
}
} catch (e: Exception) {
try { try {
if (e is InterruptedException) { withContext(realtimeContext) {
Log.d("RevoltAPI", "Socket interrupted") RealtimeSocket.connect(token)
} else {
Log.e("RevoltAPI", "WebSocket error", e)
} }
RealtimeSocket.updateDisconnectionState(DisconnectionState.Disconnected) reconnectDelay = INITIAL_RECONNECT_DELAY
} catch (e: CancellationException) {
throw e
} catch (e: SocketException) {
logcat { "WebSocket closed: ${e.message}" }
} catch (e: Exception) { } catch (e: Exception) {
logcat(LogPriority.ERROR) { "WebSocket error:\n${e.asLog()}" }
}
if (!isActive || sessionToken != token) break
try {
RealtimeSocket.updateDisconnectionState(DisconnectionState.Reconnecting)
delay(reconnectDelay)
reconnectDelay =
(reconnectDelay * 2).coerceAtMost(MAX_RECONNECT_DELAY)
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
RealtimeSocket.updateDisconnectionState(DisconnectionState.Disconnected)
Sentry.captureMessage("Error in socket error handling: $e") Sentry.captureMessage("Error in socket error handling: $e")
} }
} }
@ -214,17 +231,20 @@ object StoatAPI {
private suspend fun startSocketOps() { private suspend fun startSocketOps() {
connectWS() connectWS()
// Send a ping every roughly 30 seconds else the socket dies // Send a ping every roughly PING_INTERVAL else the socket dies
// Same interval as the web clients (/revolt.js) pingCoroutine?.cancel()
// Note: This will run even if the socket is closed (sendPing will just exit early) pingCoroutine = CoroutineScope(Dispatchers.IO).launch {
mainHandler.post(object : Runnable { while (isActive) {
override fun run() { delay(PING_INTERVAL)
runBlocking { try {
RealtimeSocket.sendPing() RealtimeSocket.sendPing()
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) { "Failed to ping WebSocket:\n${e.asLog()}" }
} }
mainHandler.postDelayed(this, 30 * 1000)
} }
}) }
} }
suspend fun initialize() { suspend fun initialize() {
@ -259,7 +279,7 @@ object StoatAPI {
unreads.clear() unreads.clear()
socketCoroutine?.cancel() socketCoroutine?.cancel()
mainHandler.removeCallbacksAndMessages(null) pingCoroutine?.cancel()
clearPersistentCache() clearPersistentCache()
} }
@ -372,4 +392,4 @@ data class RateLimitResponse(@SerialName("retry_after") val retryAfter: Int) {
internal const val NO_RETRY_AFTER = Int.MIN_VALUE internal const val NO_RETRY_AFTER = Int.MIN_VALUE
class HitRateLimitException(retryAfter: Int = NO_RETRY_AFTER) : class HitRateLimitException(retryAfter: Int = NO_RETRY_AFTER) :
Exception(if (retryAfter == NO_RETRY_AFTER) "Hit rate limit" else "Hit rate limit, retry after ${retryAfter}ms") Exception(if (retryAfter == NO_RETRY_AFTER) "Hit rate limit" else "Hit rate limit, retry after ${retryAfter}ms")

View File

@ -97,46 +97,55 @@ object RealtimeSocket {
socket?.close(CloseReason(CloseReason.Codes.NORMAL, "Reconnecting to websocket.")) socket?.close(CloseReason(CloseReason.Codes.NORMAL, "Reconnecting to websocket."))
StoatHttp.ws(STOAT_WEBSOCKET) { var activeSocket: WebSocketSession? = null
socket = this try {
StoatHttp.ws(STOAT_WEBSOCKET) {
activeSocket = this
socket = this
Log.d("RealtimeSocket", "Connected to websocket.") logcat { "Connected to websocket." }
updateDisconnectionState(DisconnectionState.Connected) updateDisconnectionState(DisconnectionState.Connected)
pushReconnectEvent() pushReconnectEvent()
// Send authorization frame // Send authorization frame
val authFrame = AuthorizationFrame("Authenticate", token) val authFrame = AuthorizationFrame("Authenticate", token)
val authFrameString = val authFrameString =
StoatJson.encodeToString(AuthorizationFrame.serializer(), authFrame) StoatJson.encodeToString(AuthorizationFrame.serializer(), authFrame)
Log.d( logcat {
"RealtimeSocket", "Sending authorization frame: ${
"Sending authorization frame: ${ authFrameString.replace(
authFrameString.replace( token,
token, "X".repeat(token.length)
"X".repeat(token.length) )
) }"
}" }
) send(StoatJson.encodeToString(AuthorizationFrame.serializer(), authFrame))
send(StoatJson.encodeToString(AuthorizationFrame.serializer(), authFrame))
incoming.consumeEach { frame -> incoming.consumeEach { frame ->
if (frame is Frame.Text) { if (frame is Frame.Text) {
val frameString = frame.readText() val frameString = frame.readText()
try { try {
val frameType = val frameType =
StoatJson.decodeFromString(AnyFrame.serializer(), frameString).type StoatJson.decodeFromString(AnyFrame.serializer(), frameString).type
handleFrame(frameType, frameString) handleFrame(frameType, frameString)
} catch (e: CancellationException) { } catch (e: CancellationException) {
throw e throw e
} catch (e: Exception) { } catch (e: Exception) {
logcat(LogPriority.ERROR) { logcat(LogPriority.ERROR) {
"Failed to handle frame: $frameString\n" + e.asLog() "Failed to handle frame: $frameString\n" + e.asLog()
}
} }
} }
} }
} }
} finally {
if (activeSocket == null || socket === activeSocket) {
socket = null
updateDisconnectionState(DisconnectionState.Disconnected)
logcat { "WebSocket disconnected." }
}
} }
} }

View File

@ -229,10 +229,6 @@ fun ChannelScreen(
val resources = LocalResources.current val resources = LocalResources.current
val config = LocalConfiguration.current val config = LocalConfiguration.current
LaunchedEffect(Unit) {
viewModel.listenToWsEvents()
}
DisposableEffect(Unit) { DisposableEffect(Unit) {
val job = scope.launch { viewModel.listenToUiCallbacks() } val job = scope.launch { viewModel.listenToUiCallbacks() }

View File

@ -40,7 +40,6 @@ import chat.stoat.api.routes.microservices.autumn.MAX_ATTACHMENTS_PER_MESSAGE
import chat.stoat.api.routes.microservices.autumn.uploadToAutumn import chat.stoat.api.routes.microservices.autumn.uploadToAutumn
import chat.stoat.api.routes.server.fetchMember import chat.stoat.api.routes.server.fetchMember
import chat.stoat.api.routes.user.addUserIfUnknown import chat.stoat.api.routes.user.addUserIfUnknown
import chat.stoat.api.routes.user.fetchUser
import chat.stoat.api.settings.GeoStateProvider import chat.stoat.api.settings.GeoStateProvider
import chat.stoat.callbacks.Action import chat.stoat.callbacks.Action
import chat.stoat.callbacks.ActionChannel import chat.stoat.callbacks.ActionChannel
@ -63,6 +62,7 @@ import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.Job import kotlinx.coroutines.Job
import kotlinx.coroutines.delay import kotlinx.coroutines.delay
import kotlinx.coroutines.flow.catch import kotlinx.coroutines.flow.catch
import kotlinx.coroutines.flow.collect
import kotlinx.coroutines.flow.first import kotlinx.coroutines.flow.first
import kotlinx.coroutines.flow.launchIn import kotlinx.coroutines.flow.launchIn
import kotlinx.coroutines.flow.onEach import kotlinx.coroutines.flow.onEach
@ -75,12 +75,18 @@ import logcat.LogPriority
import logcat.asLog import logcat.asLog
import logcat.logcat import logcat.logcat
import java.time.ZoneId import java.time.ZoneId
import kotlin.time.Duration.Companion.seconds
class ChannelScreenViewModel( class ChannelScreenViewModel(
private val kvStorage: KVStorage, private val kvStorage: KVStorage,
) : ViewModel() { ) : ViewModel() {
companion object {
private val TYPING_INDICATOR_TIMEOUT = 10.seconds
}
var items = mutableStateListOf<ChannelScreenItem>() var items = mutableStateListOf<ChannelScreenItem>()
var typingUsers = mutableStateListOf<String>() var typingUsers = mutableStateListOf<String>()
private val typingExpiryJobs = mutableMapOf<String, Job>()
var channelId by mutableStateOf<String?>(null) var channelId by mutableStateOf<String?>(null)
val channel: Channel? val channel: Channel?
@ -137,6 +143,9 @@ class ChannelScreenViewModel(
viewModelScope.launch { viewModelScope.launch {
keyboardHeight = kvStorage.getInt("keyboardHeight") ?: 900 // reasonable default for now keyboardHeight = kvStorage.getInt("keyboardHeight") ?: 900 // reasonable default for now
} }
viewModelScope.launch {
listenToWsEvents()
}
} }
private var loadMessagesJob: Job? = null private var loadMessagesJob: Job? = null
@ -145,6 +154,9 @@ class ChannelScreenViewModel(
fun switchChannel(id: String) { fun switchChannel(id: String) {
// Reset state // Reset state
this.loadMessagesJob?.cancel() this.loadMessagesJob?.cancel()
this.stopTypingJob?.cancel()
stopTyping(channelId)
clearAllTypingUsers()
requestSequence++ requestSequence++
this.channelId = id this.channelId = id
this.items = mutableStateListOf(ChannelScreenItem.Loading) this.items = mutableStateListOf(ChannelScreenItem.Loading)
@ -212,6 +224,31 @@ class ChannelScreenViewModel(
} }
} }
private fun refreshTypingUser(userId: String) {
if (userId == StoatAPI.selfId) return
if (!typingUsers.contains(userId)) {
typingUsers.add(userId)
}
typingExpiryJobs.remove(userId)?.cancel()
typingExpiryJobs[userId] = viewModelScope.launch {
delay(TYPING_INDICATOR_TIMEOUT)
typingUsers.remove(userId)
typingExpiryJobs.remove(userId)
}
}
private fun clearTypingUser(userId: String) {
typingExpiryJobs.remove(userId)?.cancel()
typingUsers.remove(userId)
}
private fun clearAllTypingUsers() {
typingExpiryJobs.values.forEach(Job::cancel)
typingExpiryJobs.clear()
typingUsers.clear()
}
private suspend fun denyMessageFieldIfNeeded() { private suspend fun denyMessageFieldIfNeeded() {
if (channel == null) return if (channel == null) return
@ -258,16 +295,19 @@ class ChannelScreenViewModel(
private fun startTyping() { private fun startTyping() {
if (editingMessage != null) return if (editingMessage != null) return
val targetChannelId = channel?.id ?: return
if (lastSentBeginTyping != null) { if (lastSentBeginTyping != null) {
val diff = Clock.System.now() - lastSentBeginTyping!! val diff = Clock.System.now() - lastSentBeginTyping!!
if (diff.inWholeSeconds < 1) return if (diff.inWholeSeconds < 1) return
} }
viewModelScope.launch { viewModelScope.launch {
withContext(StoatAPI.realtimeContext) { try {
channel?.id?.let { RealtimeSocket.beginTyping(targetChannelId)
RealtimeSocket.beginTyping(it) } catch (e: CancellationException) {
} throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) { "Failed to begin typing:\n${e.asLog()}" }
} }
} }
@ -278,18 +318,24 @@ class ChannelScreenViewModel(
private fun queueStopTyping() { private fun queueStopTyping() {
stopTypingJob = viewModelScope.launch { stopTypingJob = viewModelScope.launch {
delay(5000) delay(5.seconds)
stopTypingJob = null
stopTyping() stopTyping()
} }
} }
private fun stopTyping() { private fun stopTyping(targetChannelId: String? = channel?.id) {
lastSentBeginTyping = null
if (editingMessage != null) return if (editingMessage != null) return
if (targetChannelId == null) return
viewModelScope.launch { viewModelScope.launch {
withContext(StoatAPI.realtimeContext) { try {
channel?.id?.let { RealtimeSocket.endTyping(targetChannelId)
RealtimeSocket.endTyping(it) } catch (e: CancellationException) {
} throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) { "Failed to end typing:\n${e.asLog()}" }
} }
} }
} }
@ -373,6 +419,7 @@ class ChannelScreenViewModel(
stopTypingJob?.cancel() stopTypingJob?.cancel()
queueStopTyping() queueStopTyping()
} else { } else {
stopTypingJob?.cancel()
stopTyping() stopTyping()
} }
} }
@ -785,12 +832,74 @@ class ChannelScreenViewModel(
ackChannel(channel?.id ?: return, messageId) ackChannel(channel?.id ?: return, messageId)
} }
suspend fun listenToWsEvents() { private fun hydrateIncomingMessage(message: MessageFrame, expectedChannelId: String) {
withContext(StoatAPI.realtimeContext) { val userId = message.author
StoatAPI.wsFrameChannel.onEach { val serverId = channel?.server
if (userId != null) {
viewModelScope.launch {
try {
addUserIfUnknown(userId)
if (serverId != null && !StoatAPI.members.hasMember(serverId, userId)) {
fetchMember(serverId, userId)
}
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) {
"Failed to hydrate message author:\n${e.asLog()}"
}
}
}
}
val messageId = message.id
if (messageId != null) {
viewModelScope.launch {
try {
ackChannel(expectedChannelId, messageId)
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) { "Failed to ack message:\n${e.asLog()}" }
}
}
}
if (messageId != null && message.system == null && !message.content.isNullOrBlank()) {
viewModelScope.launch {
val ast = try {
withContext(Dispatchers.Default) { parseAst(message.content) }
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) {
"Failed to parse incoming message:\n${e.asLog()}"
}
return@launch
}
if (channelId != expectedChannelId) return@launch
val index = items.indexOfFirst { item ->
item is ChannelScreenItem.RegularMessage && item.message.id == messageId
}
val current = items.getOrNull(index) as? ChannelScreenItem.RegularMessage
?: return@launch
if (current.message.content == message.content) {
items[index] = current.copy(mdAst = ast)
}
}
}
}
private suspend fun listenToWsEvents() {
StoatAPI.wsFrameChannel.onEach {
try {
when (it) { when (it) {
is MessageFrame -> { is MessageFrame -> {
if (it.channel != channel?.id) return@onEach if (it.channel != channel?.id) return@onEach
it.author?.let(::clearTypingUser)
// If we already have the message we are just catching up on the WebSocket connection. Skip // If we already have the message we are just catching up on the WebSocket connection. Skip
if (items.any { m -> (m is ChannelScreenItem.RegularMessage && m.message.id == it.id) || (m is ChannelScreenItem.SystemMessage && m.message.id == it.id) }) return@onEach if (items.any { m -> (m is ChannelScreenItem.RegularMessage && m.message.id == it.id) || (m is ChannelScreenItem.SystemMessage && m.message.id == it.id) }) return@onEach
it.id?.let { messageId -> StoatAPI.messageCache[messageId] = it } it.id?.let { messageId -> StoatAPI.messageCache[messageId] = it }
@ -800,25 +909,10 @@ class ChannelScreenViewModel(
return@onEach return@onEach
} }
it.author?.let { userId ->
if (StoatAPI.userCache[userId] == null) {
StoatAPI.userCache[userId] = fetchUser(userId)
}
}
channel?.server?.let { serverId ->
try {
it.author?.let { userId ->
fetchMember(serverId, userId)
}
} catch (e: Exception) {
Log.e("ChannelScreenViewModel", "Failed to fetch member", e)
}
}
if (didInitialChannelFetch) { // this check is so that we don't end up with a message that arrives at the same time as the initial fetch in front of the loading indicator if (didInitialChannelFetch) { // this check is so that we don't end up with a message that arrives at the same time as the initial fetch in front of the loading indicator
val newItem = when { val newItem = when {
it.system != null -> ChannelScreenItem.SystemMessage(it) it.system != null -> ChannelScreenItem.SystemMessage(it)
else -> ChannelScreenItem.RegularMessage(it, parseAst(it.content)) else -> ChannelScreenItem.RegularMessage(it, null)
} }
updateItems(listOf(newItem) + items.filter { m -> updateItems(listOf(newItem) + items.filter { m ->
if (m is ChannelScreenItem.ProspectiveMessage) { if (m is ChannelScreenItem.ProspectiveMessage) {
@ -829,7 +923,7 @@ class ChannelScreenViewModel(
}) })
} }
it.id?.let { mid -> ackMessage(mid) } hydrateIncomingMessage(it, channel?.id ?: return@onEach)
} }
is MessageDeleteFrame -> { is MessageDeleteFrame -> {
@ -949,18 +1043,26 @@ class ChannelScreenViewModel(
is ChannelStartTypingFrame -> { is ChannelStartTypingFrame -> {
if (it.id != channel?.id) return@onEach if (it.id != channel?.id) return@onEach
if (typingUsers.contains(it.user)) return@onEach
if (it.user == StoatAPI.selfId) return@onEach if (it.user == StoatAPI.selfId) return@onEach
addUserIfUnknown(it.user) refreshTypingUser(it.user)
typingUsers.add(it.user) viewModelScope.launch {
try {
addUserIfUnknown(it.user)
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
logcat(LogPriority.ERROR) {
"Failed to hydrate typing user:\n${e.asLog()}"
}
}
}
} }
is ChannelStopTypingFrame -> { is ChannelStopTypingFrame -> {
if (it.id != channel?.id) return@onEach if (it.id != channel?.id) return@onEach
if (!typingUsers.contains(it.user)) return@onEach
typingUsers.remove(it.user) clearTypingUser(it.user)
} }
is ChannelDeleteFrame -> { is ChannelDeleteFrame -> {
@ -982,14 +1084,15 @@ class ChannelScreenViewModel(
} else { } else {
loadLatest(markLastAsRead = true) loadLatest(markLastAsRead = true)
} }
typingUsers.clear() clearAllTypingUsers()
listenToWsEvents()
} }
} }
}.catch { } catch (e: CancellationException) {
Log.e("ChannelScreen", "Failed to receive WS frame", it) throw e
}.launchIn(this) } catch (e: Exception) {
} logcat(LogPriority.ERROR) { "Failed to receive WS frame:\n${e.asLog()}" }
}
}.collect()
} }
suspend fun listenToUiCallbacks() { suspend fun listenToUiCallbacks() {