Invalidate stale relay connection cycles

This commit is contained in:
Dev 2026-07-28 03:19:20 +03:00
parent 08284082e4
commit 3c7c9d64a3
3 changed files with 350 additions and 55 deletions

View File

@ -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<String, WebSocket>()
private val reconnectJobs = ConcurrentHashMap<String, Job>()
private val desiredConnected = AtomicBoolean(false)
private val connectionEpoch = AtomicLong(0L)
private val subscriptions = ConcurrentHashMap<String, Set<String>>() // relay URL -> subscription IDs
private val messageHandlers = ConcurrentHashMap<String, (NostrEvent) -> 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
)
}
}

View File

@ -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<String>()
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<String>) {
val field = NostrRelayManager::class.java.getDeclaredField("relaysList")
field.isAccessible = true
val relays = field.get(manager) as MutableList<NostrRelayManager.Relay>
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<String, WebSocket>
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/"
}
}

View File

@ -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<String, WebSocket>
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"
}
}