Refactor IMO locking and thread leak on terminate.

This commit is contained in:
Cody Henthorne
2026-09-23 16:12:22 -04:00
committed by Michelle Tang
parent a823cfc928
commit f6289bcfdd
2 changed files with 89 additions and 69 deletions
@@ -57,10 +57,10 @@ import org.whispersystems.signalservice.api.websocket.WebSocketConnectionState
import org.whispersystems.signalservice.api.websocket.WebSocketUnavailableException
import org.whispersystems.signalservice.internal.push.Envelope
import java.util.concurrent.CopyOnWriteArrayList
import java.util.concurrent.Semaphore
import java.util.concurrent.TimeUnit
import java.util.concurrent.TimeoutException
import java.util.concurrent.atomic.AtomicInteger
import java.util.concurrent.atomic.AtomicReference
import java.util.concurrent.locks.ReentrantLock
import kotlin.concurrent.withLock
import kotlin.math.round
@@ -110,23 +110,24 @@ class IncomingMessageObserver(
private val decryptionDrainedListeners: MutableList<Runnable> = CopyOnWriteArrayList()
private val lock: ReentrantLock = ReentrantLock()
private val connectionNecessarySemaphore = Semaphore(0)
private val connectionLock = ReentrantLock()
private val connectionConditionsChanged = connectionLock.newCondition()
@Volatile
private var connectionConditionsVersion: Long = 0
private var previousSystemHttpProxy: HttpProxy? = null
private val networkConnectionListener = NetworkConnectionListener(
context = context,
onNetworkLost = { isNetworkUnavailable ->
lock.withLock {
AppDependencies.libsignalNetwork.onNetworkChange()
if (isNetworkUnavailable()) {
Log.w(TAG, "Lost network connection. Resetting the drained state.")
decryptionDrained = false
authWebSocket.disconnect()
// TODO [no-more-rest] Move the connection listener to a neutral location so this isn't passed in
unauthWebSocket.disconnect()
}
connectionNecessarySemaphore.release()
AppDependencies.libsignalNetwork.onNetworkChange()
if (isNetworkUnavailable()) {
Log.w(TAG, "Lost network connection. Resetting the drained state.")
decryptionDrained = false
authWebSocket.disconnect()
// TODO [no-more-rest] Move the connection listener to a neutral location so this isn't passed in
unauthWebSocket.disconnect()
}
notifyConnectionConditionsChanged()
},
onProxySettingsChanged = { proxyInfo ->
val systemHttpProxy = proxyInfo.toApplicableSystemHttpProxy()
@@ -142,8 +143,8 @@ class IncomingMessageObserver(
private val messageContentProcessor = MessageContentProcessor.create(context)
private var appVisible = false
private var lastInteractionTime: Long = System.currentTimeMillis()
private val appState = AtomicReference(AppState(false, System.currentTimeMillis()))
private var webSocketStateDisposable = Disposable.disposed()
private val clockSkewScope = CoroutineScope(Dispatchers.Default)
@@ -154,6 +155,20 @@ class IncomingMessageObserver(
var decryptionDrained = false
private set
private val appForegroundListener = object : AppForegroundObserver.Listener {
override fun onForeground() {
SignalExecutors.BOUNDED.execute { onAppForegrounded() }
}
override fun onBackground() {
SignalExecutors.BOUNDED.execute { onAppBackgrounded() }
}
}
private val keepAliveChangeListener: () -> Unit = {
SignalExecutors.BOUNDED.execute { notifyConnectionConditionsChanged() }
}
init {
if (INSTANCE_COUNT.incrementAndGet() != 1) {
throw AssertionError("Multiple observers!")
@@ -170,15 +185,7 @@ class IncomingMessageObserver(
}
}
AppForegroundObserver.addListener(object : AppForegroundObserver.Listener {
override fun onForeground() {
SignalExecutors.BOUNDED.execute { onAppForegrounded() }
}
override fun onBackground() {
SignalExecutors.BOUNDED.execute { onAppBackgrounded() }
}
})
AppForegroundObserver.addListener(appForegroundListener)
networkConnectionListener.register()
@@ -187,35 +194,25 @@ class IncomingMessageObserver(
.observeOn(Schedulers.computation())
.subscribeBy {
if (it == WebSocketConnectionState.CONNECTED) {
lock.withLock {
connectionNecessarySemaphore.release()
}
notifyConnectionConditionsChanged()
}
}
authWebSocket.addKeepAliveChangeListener {
SignalExecutors.BOUNDED.execute {
lock.withLock {
connectionNecessarySemaphore.release()
}
}
}
authWebSocket.addKeepAliveChangeListener(keepAliveChangeListener)
clockSkewScope.launch {
ClockSkewDetector.detected.collect { detected ->
lock.withLock {
if (detected) {
Log.w(TAG, "Clock skew detected. Disconnecting.")
authWebSocket.disconnect()
}
connectionNecessarySemaphore.release()
if (detected) {
Log.w(TAG, "Clock skew detected. Disconnecting.")
authWebSocket.disconnect()
}
notifyConnectionConditionsChanged()
}
}
}
fun notifyRegistrationStateChanged() {
connectionNecessarySemaphore.release()
notifyConnectionConditionsChanged()
}
fun notifyRestoreDecisionMade() {
@@ -235,31 +232,23 @@ class IncomingMessageObserver(
}
private fun onAppForegrounded() {
lock.withLock {
appVisible = true
ClockSkewDetector.recheck()
BackgroundService.start(context)
connectionNecessarySemaphore.release()
}
appState.updateAndGet { it.copy(isForeground = true) }
ClockSkewDetector.recheck()
BackgroundService.start(context)
notifyConnectionConditionsChanged()
}
private fun onAppBackgrounded() {
lock.withLock {
appVisible = false
ClockSkewDetector.recheck()
lastInteractionTime = System.currentTimeMillis()
connectionNecessarySemaphore.release()
}
appState.updateAndGet { it.copy(isForeground = false, lastInteractionTime = System.currentTimeMillis()) }
ClockSkewDetector.recheck()
notifyConnectionConditionsChanged()
}
private fun isConnectionNecessary(): Boolean {
val timeIdle: Long
val appVisibleSnapshot: Boolean
lock.withLock {
appVisibleSnapshot = appVisible
timeIdle = if (appVisibleSnapshot) 0 else System.currentTimeMillis() - lastInteractionTime
}
val appStateSnapshot = appState.get()
val isForeground = appStateSnapshot.isForeground
val lastInteractionTime = appStateSnapshot.lastInteractionTime
val timeIdle = if (isForeground) 0 else System.currentTimeMillis() - lastInteractionTime
val registered = SignalStore.account.isRegistered
val unauthorizedReceived = TextSecurePreferences.isUnauthorizedReceived(context)
@@ -271,18 +260,18 @@ class IncomingMessageObserver(
val clockSkewDetected = ClockSkewDetector.isDetected
val clockSkew = ClockSkewDetector.skew
val lastInteractionString = if (appVisibleSnapshot) "N/A" else timeIdle.toString() + " ms (" + (if (timeIdle < maxBackgroundTime) "within limit" else "over limit") + ")"
val lastInteractionString = if (isForeground) "N/A" else timeIdle.toString() + " ms (" + (if (timeIdle < maxBackgroundTime) "within limit" else "over limit") + ")"
val conclusion = registered &&
!unauthorizedReceived &&
!clockSkewDetected &&
(appVisibleSnapshot || timeIdle < maxBackgroundTime || !fcmEnabled || forceWebsocket) &&
(isForeground || timeIdle < maxBackgroundTime || !fcmEnabled || forceWebsocket) &&
hasNetwork
val needsConnectionString = if (conclusion) "Needs Connection" else "Does Not Need Connection"
Log.d(
TAG,
"[$needsConnectionString] Network: $hasNetwork, Foreground: $appVisibleSnapshot, Time Since Last Interaction: $lastInteractionString, FCM: $fcmEnabled, WS Open or Keep-alives: $websocketAlreadyOpen, Registered: $registered, Unauthorized: $unauthorizedReceived, Proxy: $hasProxy, Force websocket: $forceWebsocket, Clock skew: $clockSkewDetected ($clockSkew)"
"[$needsConnectionString] Network: $hasNetwork, Foreground: $isForeground, Time Since Last Interaction: $lastInteractionString, FCM: $fcmEnabled, WS Open or Keep-alives: $websocketAlreadyOpen, Registered: $registered, Unauthorized: $unauthorizedReceived, Proxy: $hasProxy, Force websocket: $forceWebsocket, Clock skew: $clockSkewDetected ($clockSkew)"
)
return conclusion
@@ -292,13 +281,21 @@ class IncomingMessageObserver(
return !TextSecurePreferences.isUnauthorizedReceived(context) && SignalStore.account.isRegistered && (authWebSocket.stateSnapshot == WebSocketConnectionState.CONNECTED || (authWebSocket.shouldSendKeepAlives() && NetworkConstraint.isMet(context)))
}
/** Conditions are evaluated outside of [connectionLock] so notifiers never block on them. We only park if nothing signaled since the version we evaluated against. */
private fun waitForConnectionNecessary() {
try {
connectionNecessarySemaphore.drainPermits()
while (ClockSkewDetector.isDetected || (!isConnectionNecessary() && !isConnectionAvailable())) {
val numberDrained = connectionNecessarySemaphore.drainPermits()
if (numberDrained == 0) {
connectionNecessarySemaphore.acquire()
var observedVersion = connectionConditionsVersion
while (!terminated) {
if (!ClockSkewDetector.isDetected && (isConnectionNecessary() || isConnectionAvailable())) {
return
}
connectionLock.withLock {
if (!terminated && connectionConditionsVersion == observedVersion) {
connectionConditionsChanged.await()
}
observedVersion = connectionConditionsVersion
}
}
} catch (e: InterruptedException) {
@@ -306,15 +303,25 @@ class IncomingMessageObserver(
}
}
private fun notifyConnectionConditionsChanged() {
connectionLock.withLock {
connectionConditionsVersion++
connectionConditionsChanged.signalAll()
}
}
fun terminate() {
Log.w(TAG, "Termination! ${this.hashCode()}", Throwable())
INSTANCE_COUNT.decrementAndGet()
AppForegroundObserver.removeListener(appForegroundListener)
authWebSocket.removeKeepAliveChangeListener(keepAliveChangeListener)
networkConnectionListener.unregister()
webSocketStateDisposable.dispose()
clockSkewScope.cancel()
terminated = true
authWebSocket.removeKeepAliveToken(WEB_SOCKET_KEEP_ALIVE_TOKEN)
authWebSocket.disconnect()
notifyConnectionConditionsChanged()
}
@VisibleForTesting
@@ -476,6 +483,10 @@ class IncomingMessageObserver(
}
waitForConnectionNecessary()
if (terminated) {
break
}
Log.i(TAG, "Making websocket connection....")
val webSocketDisposable = authWebSocket.state.subscribe { state: WebSocketConnectionState ->
@@ -559,7 +570,7 @@ class IncomingMessageObserver(
}
}
if (!appVisible) {
if (!appState.get().isForeground) {
BackgroundService.stop(context)
}
} catch (e: Throwable) {
@@ -749,6 +760,11 @@ class IncomingMessageObserver(
}
}
private data class AppState(
val isForeground: Boolean,
val lastInteractionTime: Long
)
data class ProcessingResult(
val followUpOperations: List<FollowUpOperation>,
val isNetworkResetRequired: Boolean = false
@@ -167,6 +167,10 @@ sealed class SignalWebSocket(
keepAliveChangeListeners.add(listener)
}
fun removeKeepAliveChangeListener(listener: Listener) {
keepAliveChangeListeners.remove(listener)
}
fun request(request: WebSocketRequestMessage): Single<WebsocketResponse> {
return try {
restartDelayedDisconnectIfNecessary()