From 3c7c9d64a3a5e43a44af660b38e4eebd7375109b Mon Sep 17 00:00:00 2001 From: Dev Date: Tue, 28 Jul 2026 03:19:20 +0300 Subject: [PATCH] Invalidate stale relay connection cycles --- .../android/nostr/NostrRelayManager.kt | 168 ++++++++++++------ .../nostr/NostrRelayConnectionEpochTest.kt | 141 +++++++++++++++ .../nostr/NostrRelayResetCallbackTest.kt | 96 ++++++++++ 3 files changed, 350 insertions(+), 55 deletions(-) create mode 100644 app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayConnectionEpochTest.kt create mode 100644 app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayResetCallbackTest.kt diff --git a/app/src/main/java/com/bitchat/android/nostr/NostrRelayManager.kt b/app/src/main/java/com/bitchat/android/nostr/NostrRelayManager.kt index 62266af6..cb88787c 100644 --- a/app/src/main/java/com/bitchat/android/nostr/NostrRelayManager.kt +++ b/app/src/main/java/com/bitchat/android/nostr/NostrRelayManager.kt @@ -33,7 +33,8 @@ class NostrRelayManager internal constructor( private val scope: CoroutineScope = CoroutineScope(Dispatchers.IO + SupervisorJob()), private val eventDeduplicator: NostrEventDeduplicator = - NostrEventDeduplicator.getInstance() + NostrEventDeduplicator.getInstance(), + private val webSocketFactory: ((Request, WebSocketListener) -> WebSocket)? = null ) { companion object { @@ -98,6 +99,7 @@ class NostrRelayManager internal constructor( private val connections = ConcurrentHashMap() private val reconnectJobs = ConcurrentHashMap() private val desiredConnected = AtomicBoolean(false) + private val connectionEpoch = AtomicLong(0L) private val subscriptions = ConcurrentHashMap>() // relay URL -> subscription IDs private val messageHandlers = ConcurrentHashMap Unit>() private val commitAwareMessageHandlers = @@ -333,6 +335,9 @@ class NostrRelayManager internal constructor( LiveLocationPrivacyGate.runIfAllowed(liveLocationToken, action) } + private fun isCurrentConnectionEpoch(epoch: Long): Boolean = + desiredConnected.get() && connectionEpoch.get() == epoch + /** * Privacy teardown is allowed to bypass an already-revoked token solely to stop * server-side delivery. Live subscription IDs are opaque, so CLOSE carries no @@ -372,11 +377,13 @@ class NostrRelayManager internal constructor( relayUrl in nonLiveRelayUrls && isCurrentAccountGeneration(generation) ) { + val epoch = connectionEpoch.get() scope.launch { connectToRelay( relayUrl, liveLocationToken = null, - generation = generation + generation = generation, + expectedConnectionEpoch = epoch ) } } @@ -462,20 +469,26 @@ class NostrRelayManager internal constructor( if (!tracked) return updateRelaysList() - if (!desiredConnected.get()) return + val epoch = connectionEpoch.get() + if (!isCurrentConnectionEpoch(epoch)) return val job = scope.launch { - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(epoch) || !isCurrentAccountGeneration(generation) || !isNetworkActionAllowed(liveLocationToken) ) return@launch relayUrls.forEach { relayUrl -> launch { - if (desiredConnected.get() && + if (isCurrentConnectionEpoch(epoch) && isCurrentAccountGeneration(generation) && !connections.containsKey(relayUrl) && isNetworkActionAllowed(liveLocationToken) ) { - connectToRelay(relayUrl, liveLocationToken, generation) + connectToRelay( + relayUrl, + liveLocationToken, + generation, + epoch + ) } } } @@ -519,35 +532,35 @@ class NostrRelayManager internal constructor( * Connect to all configured relays */ fun connect() { - val (generation, relayUrls) = synchronized(accountGenerationLock) { + val (generation, epoch, relayUrls) = synchronized(accountGenerationLock) { val current = accountGeneration.get() if (!isCurrentAccountGeneration(current)) return desiredConnected.set(true) - current to synchronized(relaysList) { + Triple(current, connectionEpoch.get(), synchronized(relaysList) { relaysList.map { it.url } - } + }) } Log.i(TAG, "Connecting to ${relayUrls.size} Nostr relays") scope.launch { - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(epoch) || !isCurrentAccountGeneration(generation) ) return@launch relayUrls.forEach { relayUrl -> launch { val liveToken = liveLocationRelayTokens[relayUrl] ?.takeIf { relayUrl !in nonLiveRelayUrls } - if (desiredConnected.get() && + if (isCurrentConnectionEpoch(epoch) && isCurrentAccountGeneration(generation) && (liveToken == null || LiveLocationPrivacyGate.accepts(liveToken)) ) { - connectToRelay(relayUrl, liveToken, generation) + connectToRelay(relayUrl, liveToken, generation, epoch) } } } } // Start periodic subscription validation - startSubscriptionValidation(generation) + startSubscriptionValidation(generation, epoch) } /** @@ -555,7 +568,10 @@ class NostrRelayManager internal constructor( */ fun disconnect() { Log.i(TAG, "Disconnecting from all Nostr relays") - desiredConnected.set(false) + synchronized(accountGenerationLock) { + desiredConnected.set(false) + connectionEpoch.incrementAndGet() + } // Stop subscription validation stopSubscriptionValidation() @@ -975,9 +991,10 @@ class NostrRelayManager internal constructor( val generation = accountGeneration.get() if (!isCurrentAccountGeneration(generation)) return val relay = relaysList.find { it.url == relayUrl } ?: return - synchronized(accountGenerationLock) { + val epoch = synchronized(accountGenerationLock) { if (!isCurrentAccountGeneration(generation)) return desiredConnected.set(true) + connectionEpoch.get() } val liveToken = liveLocationRelayTokens[relayUrl] ?.takeIf { relayUrl !in nonLiveRelayUrls } @@ -993,10 +1010,10 @@ class NostrRelayManager internal constructor( // Attempt immediate reconnection scope.launch { - if (desiredConnected.get() && + if (isCurrentConnectionEpoch(epoch) && isCurrentAccountGeneration(generation) ) { - connectToRelay(relayUrl, liveToken, generation) + connectToRelay(relayUrl, liveToken, generation, epoch) } } } @@ -1124,6 +1141,7 @@ class NostrRelayManager internal constructor( ) return@generationCheck null desiredConnected.set(false) + connectionEpoch.incrementAndGet() val jobs = buildList { subscriptionValidationJob?.let(::add) addAll(reconnectJobs.values) @@ -1318,18 +1336,26 @@ class NostrRelayManager internal constructor( /** * Start periodic subscription validation to ensure robustness */ - private fun startSubscriptionValidation(generation: Long) { - if (!isCurrentAccountGeneration(generation)) return + private fun startSubscriptionValidation( + generation: Long, + connectionEpoch: Long + ) { + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) return stopSubscriptionValidation() // Stop any existing validation subscriptionValidationJob = scope.launch { - if (!isCurrentAccountGeneration(generation)) return@launch + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) return@launch val manager = powerManager if (manager == null) { runSubscriptionValidationLoop( intervalMs = com.bitchat.android.util.AppConstants.Nostr .SUBSCRIPTION_VALIDATION_INTERVAL_MS, - generation = generation + generation = generation, + connectionEpoch = connectionEpoch ) return@launch } @@ -1338,36 +1364,48 @@ class NostrRelayManager internal constructor( .map { it.nostr.subscriptionValidationMs } .distinctUntilChanged() .collectLatest { intervalMs -> - runSubscriptionValidationLoop(intervalMs, generation) + runSubscriptionValidationLoop( + intervalMs, + generation, + connectionEpoch + ) } } } private suspend fun runSubscriptionValidationLoop( intervalMs: Long, - generation: Long + generation: Long, + connectionEpoch: Long ) { while (currentCoroutineContext().isActive && - desiredConnected.get() && + isCurrentConnectionEpoch(connectionEpoch) && isCurrentAccountGeneration(generation) ) { delay(intervalMs) - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(connectionEpoch) || !isCurrentAccountGeneration(generation) ) break - validateAndRepairSubscriptions(generation) + validateAndRepairSubscriptions(generation, connectionEpoch) } } - private fun validateAndRepairSubscriptions(generation: Long) { - if (!isCurrentAccountGeneration(generation)) return + private fun validateAndRepairSubscriptions( + generation: Long, + connectionEpoch: Long + ) { + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) return try { val report = validateSubscriptionConsistency() if (report.isConsistent || report.connectedRelayCount == 0) return Log.w(TAG, "Nostr subscription inconsistencies detected") connections.forEach { (relayUrl, webSocket) -> - if (!isCurrentAccountGeneration(generation)) return + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) return val currentSubs = subscriptions[relayUrl] ?: emptySet() val expectedSubs = activeSubscriptions.keys.filter { subId -> val subInfo = activeSubscriptions[subId] @@ -1405,11 +1443,12 @@ class NostrRelayManager internal constructor( private suspend fun connectToRelay( urlString: String, liveLocationToken: Long? = null, - generation: Long = accountGeneration.get() + generation: Long = accountGeneration.get(), + expectedConnectionEpoch: Long = connectionEpoch.get() ) { val connectionToken = liveLocationToken ?.takeIf { urlString !in nonLiveRelayUrls } - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(expectedConnectionEpoch) || !isCurrentAccountGeneration(generation) || !isNetworkActionAllowed(connectionToken) ) return @@ -1424,18 +1463,18 @@ class NostrRelayManager internal constructor( .build() val started = runNetworkAction(connectionToken) { - val webSocket = httpClient.newWebSocket( - request, - RelayWebSocketListener( - relayUrl = urlString, - liveLocationToken = connectionToken, - generation = generation - ) + val listener = RelayWebSocketListener( + relayUrl = urlString, + liveLocationToken = connectionToken, + generation = generation, + connectionEpoch = expectedConnectionEpoch ) + val webSocket = webSocketFactory?.invoke(request, listener) + ?: httpClient.newWebSocket(request, listener) val existing = connections.putIfAbsent(urlString, webSocket) when { existing != null -> webSocket.close(1000, "Duplicate connection") - !desiredConnected.get() || + !isCurrentConnectionEpoch(expectedConnectionEpoch) || !isCurrentAccountGeneration(generation) || !isNetworkActionAllowed(connectionToken) -> { connections.remove(urlString, webSocket) @@ -1451,7 +1490,8 @@ class NostrRelayManager internal constructor( relayUrl = urlString, error = e, liveLocationToken = connectionToken, - generation = generation + generation = generation, + connectionEpoch = expectedConnectionEpoch ) } } @@ -1638,9 +1678,12 @@ class NostrRelayManager internal constructor( webSocket: WebSocket, error: Throwable, liveLocationToken: Long? = null, - generation: Long + generation: Long, + connectionEpoch: Long ) { - if (!isCurrentAccountGeneration(generation)) { + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) { connections.remove(relayUrl, webSocket) return } @@ -1652,7 +1695,8 @@ class NostrRelayManager internal constructor( relayUrl, error, liveLocationToken, - generation + generation, + connectionEpoch ) } @@ -1660,9 +1704,10 @@ class NostrRelayManager internal constructor( relayUrl: String, error: Throwable, liveLocationToken: Long?, - generation: Long + generation: Long, + connectionEpoch: Long ) { - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(connectionEpoch) || !isCurrentAccountGeneration(generation) || connections.containsKey(relayUrl) ) return @@ -1670,7 +1715,8 @@ class NostrRelayManager internal constructor( relayUrl, error, liveLocationToken, - generation + generation, + connectionEpoch ) } @@ -1678,9 +1724,12 @@ class NostrRelayManager internal constructor( relayUrl: String, error: Throwable, liveLocationToken: Long?, - generation: Long + generation: Long, + connectionEpoch: Long ) { - if (!isCurrentAccountGeneration(generation)) return + if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) + ) return val connectionToken = liveLocationToken ?.takeIf { relayUrl !in nonLiveRelayUrls } @@ -1697,7 +1746,7 @@ class NostrRelayManager internal constructor( ) } } - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(connectionEpoch) || !isCurrentAccountGeneration(generation) || !isNetworkActionAllowed(connectionToken) ) return @@ -1738,11 +1787,16 @@ class NostrRelayManager internal constructor( reconnectJobs.remove(relayUrl)?.cancel() val reconnectJob = scope.launch { delay(backoffInterval) - if (desiredConnected.get() && + if (isCurrentConnectionEpoch(connectionEpoch) && isCurrentAccountGeneration(generation) && isNetworkActionAllowed(connectionToken) ) { - connectToRelay(relayUrl, connectionToken, generation) + connectToRelay( + relayUrl, + connectionToken, + generation, + connectionEpoch + ) } } reconnectJobs[relayUrl] = reconnectJob @@ -1841,11 +1895,12 @@ class NostrRelayManager internal constructor( private inner class RelayWebSocketListener( private val relayUrl: String, private val liveLocationToken: Long?, - private val generation: Long + private val generation: Long, + private val connectionEpoch: Long ) : WebSocketListener() { override fun onOpen(webSocket: WebSocket, response: Response) { - if (!desiredConnected.get() || + if (!isCurrentConnectionEpoch(connectionEpoch) || !isCurrentAccountGeneration(generation) || connections[relayUrl] !== webSocket || !isNetworkActionAllowed(liveLocationToken) @@ -1895,6 +1950,7 @@ class NostrRelayManager internal constructor( override fun onMessage(webSocket: WebSocket, text: String) { if (!isCurrentAccountGeneration(generation) || + !isCurrentConnectionEpoch(connectionEpoch) || connections[relayUrl] !== webSocket ) return handleMessage(text, relayUrl, generation) @@ -1911,7 +1967,8 @@ class NostrRelayManager internal constructor( webSocket, error, liveLocationToken, - generation + generation, + connectionEpoch ) } @@ -1922,7 +1979,8 @@ class NostrRelayManager internal constructor( webSocket, t, liveLocationToken, - generation + generation, + connectionEpoch ) } } diff --git a/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayConnectionEpochTest.kt b/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayConnectionEpochTest.kt new file mode 100644 index 00000000..83277eea --- /dev/null +++ b/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayConnectionEpochTest.kt @@ -0,0 +1,141 @@ +package com.bitchat.android.nostr + +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.cancel +import kotlinx.coroutines.test.StandardTestDispatcher +import okhttp3.Request +import okhttp3.WebSocket +import okhttp3.WebSocketListener +import okio.ByteString +import org.junit.Assert.assertEquals +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner + +@RunWith(RobolectricTestRunner::class) +class NostrRelayConnectionEpochTest { + @Test + fun `disconnect invalidates already queued connect work before reconnect`() { + val dispatcher = StandardTestDispatcher() + val scope = CoroutineScope(dispatcher + SupervisorJob()) + val openedUrls = mutableListOf() + val manager = NostrRelayManager( + scope = scope, + eventDeduplicator = NostrEventDeduplicator(maxCapacity = 8), + webSocketFactory = { request, _ -> + openedUrls += request.url.host + RecordingWebSocket(request) + } + ) + + try { + manager.connect() + manager.disconnect() + replaceRelays(manager, listOf(FRESH_RELAY_URL)) + manager.connect() + + dispatcher.scheduler.runCurrent() + + assertEquals(listOf(FRESH_RELAY_HOST), openedUrls) + } finally { + manager.disconnect() + scope.cancel() + } + } + + @Test + fun `message callback from replaced socket cannot enter a current subscription`() { + val scope = CoroutineScope(Dispatchers.Unconfined + SupervisorJob()) + val deduplicator = NostrEventDeduplicator(maxCapacity = 8) + var listener: WebSocketListener? = null + lateinit var originalSocket: RecordingWebSocket + val manager = NostrRelayManager( + scope = scope, + eventDeduplicator = deduplicator, + webSocketFactory = { request, createdListener -> + listener = createdListener + RecordingWebSocket(request).also { originalSocket = it } + } + ) + replaceRelays(manager, listOf(FRESH_RELAY_URL)) + var processed = 0 + val event = NostrEvent( + id = "7a".repeat(32), + pubkey = "7b".repeat(32), + createdAt = 1, + kind = 1060, + tags = emptyList(), + content = "ciphertext", + sig = "signature" + ) + + try { + manager.subscribeAfterSuccessfulProcessing( + filter = NostrFilter(kinds = listOf(1060)), + id = "current-subscription", + targetRelayUrls = listOf(FRESH_RELAY_URL) + ) { + processed += 1 + true + } + manager.connect() + installConnection( + manager, + FRESH_RELAY_URL, + RecordingWebSocket(Request.Builder().url(FRESH_RELAY_URL).build()) + ) + + requireNotNull(listener).onMessage( + originalSocket, + """["EVENT","current-subscription",${event.toJsonString()}]""" + ) + + assertEquals(0, processed) + assertEquals(false, deduplicator.contains(event.id)) + } finally { + manager.disconnect() + scope.cancel() + } + } + + @Suppress("UNCHECKED_CAST") + private fun replaceRelays(manager: NostrRelayManager, urls: List) { + val field = NostrRelayManager::class.java.getDeclaredField("relaysList") + field.isAccessible = true + val relays = field.get(manager) as MutableList + synchronized(relays) { + relays.clear() + relays.addAll(urls.map(NostrRelayManager::Relay)) + } + } + + @Suppress("UNCHECKED_CAST") + private fun installConnection( + manager: NostrRelayManager, + relayUrl: String, + socket: WebSocket + ) { + val field = NostrRelayManager::class.java.getDeclaredField("connections") + field.isAccessible = true + val connections = field.get(manager) as MutableMap + connections[relayUrl] = socket + } + + private class RecordingWebSocket( + private val request: Request + ) : WebSocket { + override fun request(): Request = request + override fun queueSize(): Long = 0L + override fun send(text: String): Boolean = true + override fun send(bytes: ByteString): Boolean = true + override fun close(code: Int, reason: String?): Boolean = true + override fun cancel() = Unit + } + + companion object { + private const val FRESH_RELAY_HOST = "fresh-cycle.example" + private const val FRESH_RELAY_URL = "wss://$FRESH_RELAY_HOST/" + } +} diff --git a/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayResetCallbackTest.kt b/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayResetCallbackTest.kt new file mode 100644 index 00000000..1c1c60ca --- /dev/null +++ b/app/src/test/kotlin/com/bitchat/android/nostr/NostrRelayResetCallbackTest.kt @@ -0,0 +1,96 @@ +package com.bitchat.android.nostr + +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.cancel +import okhttp3.Request +import okhttp3.WebSocket +import okio.ByteString +import org.junit.Assert.assertTrue +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner +import java.util.concurrent.CountDownLatch +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicBoolean +import kotlin.concurrent.thread + +@RunWith(RobolectricTestRunner::class) +class NostrRelayResetCallbackTest { + @Test + fun `reset confirmation callback runs outside account locks`() { + val scope = CoroutineScope(Dispatchers.Unconfined + SupervisorJob()) + val manager = NostrRelayManager( + scope = scope, + eventDeduplicator = NostrEventDeduplicator(maxCapacity = 8) + ) + installConnection(manager, RELAY_URL, RecordingWebSocket()) + val callbackReachedAccountLock = AtomicBoolean(false) + val event = signedEvent() + manager.sendEventConfirmed(event, listOf(RELAY_URL)) { + val lockReached = CountDownLatch(1) + thread(start = true, name = "relay-reset-callback-lock-probe") { + manager.registerPendingGiftWrap( + id = "callback-probe", + expectedAccountGeneration = manager.captureAccountGeneration() + ) + lockReached.countDown() + } + callbackReachedAccountLock.set(lockReached.await(1, TimeUnit.SECONDS)) + } + + val resetToken = manager.beginAccountReset() + val resetFinished = CountDownLatch(1) + thread(start = true, name = "relay-reset-probe") { + manager.discardForAccountReset(resetToken) + resetFinished.countDown() + } + + try { + assertTrue(resetFinished.await(2, TimeUnit.SECONDS)) + assertTrue(callbackReachedAccountLock.get()) + assertTrue(manager.completeAccountReset(resetToken)) + } finally { + scope.cancel() + } + } + + @Suppress("UNCHECKED_CAST") + private fun installConnection( + manager: NostrRelayManager, + relayUrl: String, + socket: WebSocket + ) { + val field = NostrRelayManager::class.java.getDeclaredField("connections") + field.isAccessible = true + val connections = field.get(manager) as MutableMap + connections[relayUrl] = socket + } + + private fun signedEvent(): NostrEvent { + val privateKey = "0".repeat(63) + "1" + return NostrEvent( + pubkey = NostrCrypto.derivePublicKey(privateKey), + createdAt = 1, + kind = 1060, + tags = emptyList(), + content = "ciphertext" + ).sign(privateKey) + } + + private class RecordingWebSocket : WebSocket { + override fun request(): Request = + Request.Builder().url(RELAY_URL).build() + + override fun queueSize(): Long = 0L + override fun send(text: String): Boolean = true + override fun send(bytes: ByteString): Boolean = true + override fun close(code: Int, reason: String?): Boolean = true + override fun cancel() = Unit + } + + companion object { + private const val RELAY_URL = "wss://relay-reset.example" + } +}