From b87239cc86509ff567d8c676c431cd98ab30d2b7 Mon Sep 17 00:00:00 2001 From: jack Date: Sun, 12 Jul 2026 11:15:09 -0400 Subject: [PATCH] Make Wi-Fi peer rebinding atomic --- .../wifi-aware/WifiAwareConnectionTracker.kt | 80 +++++++++++++----- .../wifi-aware/WifiAwareMeshService.kt | 8 +- .../WifiAwareConnectionTrackerTest.kt | 81 +++++++++++++++++++ docs/NOISE_PEER_ID_BINDING.md | 5 +- 4 files changed, 151 insertions(+), 23 deletions(-) create mode 100644 app/src/test/kotlin/com/bitchat/android/wifi-aware/WifiAwareConnectionTrackerTest.kt diff --git a/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareConnectionTracker.kt b/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareConnectionTracker.kt index 115dbb05..9376b078 100644 --- a/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareConnectionTracker.kt +++ b/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareConnectionTracker.kt @@ -23,6 +23,7 @@ class WifiAwareConnectionTracker( // Active resources per peer val peerSockets = ConcurrentHashMap() private val socketAliases = ConcurrentHashMap() + private val socketBindingLock = Any() val serverSockets = ConcurrentHashMap() val networkCallbacks = ConcurrentHashMap() @@ -32,24 +33,26 @@ class WifiAwareConnectionTracker( } override fun disconnect(id: String) { - Log.d(TAG, "Disconnecting peer $id") - val canonicalId = resolveCanonicalPeerId(id) - - // 1. Close client socket - peerSockets.remove(canonicalId)?.let { - try { it.close() } catch (e: Exception) { Log.w(TAG, "Error closing socket for $id: ${e.message}") } - } - socketAliases.entries.removeIf { it.key == id || it.key == canonicalId || it.value == canonicalId } + synchronized(socketBindingLock) { + Log.d(TAG, "Disconnecting peer $id") + val canonicalId = resolveCanonicalPeerId(id) - // 2. Close server socket - serverSockets.remove(canonicalId)?.let { - try { it.close() } catch (e: Exception) { Log.w(TAG, "Error closing server socket for $id: ${e.message}") } - } + // 1. Close client socket + peerSockets.remove(canonicalId)?.let { + try { it.close() } catch (e: Exception) { Log.w(TAG, "Error closing socket for $id: ${e.message}") } + } + socketAliases.entries.removeIf { it.key == id || it.key == canonicalId || it.value == canonicalId } - // Ensure any pending/active network request is explicitly released - releaseNetworkRequest(canonicalId) - removePendingConnection(id) - removePendingConnection(canonicalId) + // 2. Close server socket + serverSockets.remove(canonicalId)?.let { + try { it.close() } catch (e: Exception) { Log.w(TAG, "Error closing server socket for $id: ${e.message}") } + } + + // Ensure any pending/active network request is explicitly released + releaseNetworkRequest(canonicalId) + removePendingConnection(id) + removePendingConnection(canonicalId) + } } fun releaseNetworkRequest(id: String) { @@ -71,12 +74,16 @@ class WifiAwareConnectionTracker( * Successfully established a client connection */ fun onClientConnected(peerId: String, socket: SyncedSocket) { - // Close previous socket if one exists to prevent zombie readers - peerSockets[peerId]?.let { - try { it.close() } catch (_: Exception) {} + synchronized(socketBindingLock) { + val canonicalPeerId = resolveCanonicalPeerId(peerId) + // Close previous socket if one exists to prevent zombie readers + peerSockets[canonicalPeerId]?.let { + try { it.close() } catch (_: Exception) {} + } + peerSockets[canonicalPeerId] = socket + removePendingConnection(peerId) // Clear retry state on success + if (canonicalPeerId != peerId) removePendingConnection(canonicalPeerId) } - peerSockets[peerId] = socket - removePendingConnection(peerId) // Clear retry state on success } fun getSocketForPeer(peerId: String): SyncedSocket? { @@ -87,6 +94,37 @@ class WifiAwareConnectionTracker( fun canonicalPeerId(peerId: String): String = resolveCanonicalPeerId(peerId) fun rebindPeerId(previousPeerId: String, resolvedPeerId: String, socket: SyncedSocket): String { + return synchronized(socketBindingLock) { + rebindPeerIdLocked(previousPeerId, resolvedPeerId, socket) + } + } + + /** + * Atomically require that [expectedSocket] is still the active provisional transport and, only + * then, promote it. This closes the gap where a replacement socket could land after validation + * but before mutation and the stale authenticated socket would become canonical. + */ + fun rebindPeerIdIfCurrent( + previousPeerId: String, + resolvedPeerId: String, + expectedSocket: SyncedSocket + ): Boolean = synchronized(socketBindingLock) { + val previousCanonical = resolveCanonicalPeerId(previousPeerId) + if (peerSockets[previousCanonical] !== expectedSocket) return@synchronized false + val resolvedCanonical = resolveCanonicalPeerId(resolvedPeerId) + val existingResolvedSocket = peerSockets[resolvedCanonical] + if (existingResolvedSocket != null && existingResolvedSocket !== expectedSocket) { + return@synchronized false + } + rebindPeerIdLocked(previousPeerId, resolvedPeerId, expectedSocket) + true + } + + private fun rebindPeerIdLocked( + previousPeerId: String, + resolvedPeerId: String, + socket: SyncedSocket + ): String { if (previousPeerId == resolvedPeerId) { peerSockets[resolvedPeerId] = socket return resolvedPeerId diff --git a/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareMeshService.kt b/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareMeshService.kt index 9956b84a..e463f5f8 100644 --- a/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareMeshService.kt +++ b/app/src/main/java/com/bitchat/android/wifi-aware/WifiAwareMeshService.kt @@ -1291,7 +1291,13 @@ class WifiAwareMeshService(private val context: Context) : MeshService, Transpor return } - connectionTracker.rebindPeerId(provisionalPeerId, canonicalPeerId, link.transport) + if (!connectionTracker.rebindPeerIdIfCurrent(provisionalPeerId, canonicalPeerId, link.transport)) { + Log.w( + TAG, + "Ignoring Noise link promotion for ${canonicalPeerId.take(8)}: provisional socket changed before rebind" + ) + return + } handleToPeerId.forEach { (handle, peerId) -> if (peerId == provisionalPeerId) handleToPeerId[handle] = canonicalPeerId } diff --git a/app/src/test/kotlin/com/bitchat/android/wifi-aware/WifiAwareConnectionTrackerTest.kt b/app/src/test/kotlin/com/bitchat/android/wifi-aware/WifiAwareConnectionTrackerTest.kt new file mode 100644 index 00000000..b1b12d0e --- /dev/null +++ b/app/src/test/kotlin/com/bitchat/android/wifi-aware/WifiAwareConnectionTrackerTest.kt @@ -0,0 +1,81 @@ +package com.bitchat.android.wifiaware + +import android.net.ConnectivityManager +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import org.junit.Assert.assertFalse +import org.junit.Assert.assertNull +import org.junit.Assert.assertSame +import org.junit.Assert.assertTrue +import org.junit.Test +import org.mockito.kotlin.doReturn +import org.mockito.kotlin.mock +import java.io.ByteArrayInputStream +import java.io.ByteArrayOutputStream +import java.net.Socket + +class WifiAwareConnectionTrackerTest { + @Test + fun `compare and rebind rejects stale authenticated socket after replacement`() { + val tracker = WifiAwareConnectionTracker( + CoroutineScope(SupervisorJob() + Dispatchers.Unconfined), + mock() + ) + val authenticatedSocket = syncedSocket() + val replacementSocket = syncedSocket() + tracker.onClientConnected("provisional", authenticatedSocket) + tracker.onClientConnected("provisional", replacementSocket) + + assertFalse( + tracker.rebindPeerIdIfCurrent( + previousPeerId = "provisional", + resolvedPeerId = "canonical", + expectedSocket = authenticatedSocket + ) + ) + assertSame(replacementSocket, tracker.getSocketForPeer("provisional")) + assertNull(tracker.getSocketForPeer("canonical")) + + assertTrue( + tracker.rebindPeerIdIfCurrent( + previousPeerId = "provisional", + resolvedPeerId = "canonical", + expectedSocket = replacementSocket + ) + ) + assertSame(replacementSocket, tracker.getSocketForPeer("canonical")) + assertSame(replacementSocket, tracker.getSocketForPeer("provisional")) + } + + @Test + fun `authenticated provisional socket cannot displace existing canonical socket`() { + val tracker = WifiAwareConnectionTracker( + CoroutineScope(SupervisorJob() + Dispatchers.Unconfined), + mock() + ) + val provisionalSocket = syncedSocket() + val canonicalSocket = syncedSocket() + tracker.onClientConnected("provisional", provisionalSocket) + tracker.onClientConnected("canonical", canonicalSocket) + + assertFalse( + tracker.rebindPeerIdIfCurrent( + previousPeerId = "provisional", + resolvedPeerId = "canonical", + expectedSocket = provisionalSocket + ) + ) + assertSame(provisionalSocket, tracker.getSocketForPeer("provisional")) + assertSame(canonicalSocket, tracker.getSocketForPeer("canonical")) + assertTrue("Rejected promotion must not alias the provisional ID", tracker.canonicalPeerId("provisional") == "provisional") + } + + private fun syncedSocket(): SyncedSocket { + val raw = mock { + on { getInputStream() } doReturn ByteArrayInputStream(byteArrayOf()) + on { getOutputStream() } doReturn ByteArrayOutputStream() + } + return SyncedSocket(raw) + } +} diff --git a/docs/NOISE_PEER_ID_BINDING.md b/docs/NOISE_PEER_ID_BINDING.md index 95fe1b6e..1b2e319c 100644 --- a/docs/NOISE_PEER_ID_BINDING.md +++ b/docs/NOISE_PEER_ID_BINDING.md @@ -18,7 +18,10 @@ validation succeeds. A Wi-Fi discovery identity is not destructively rebound from a self-signed announce. A direct announce may start the canonical handshake, but the alias is promoted only when the exact, still-active socket delivers the Noise frame that completes bound authentication. A peer-ID-only -or stale-socket callback cannot authorize that rebind. +or stale-socket callback cannot authorize that rebind; the final +expected-socket comparison and alias mutation are atomic with socket replacement. +Promotion also refuses to displace a different live socket already authenticated +under the canonical peer ID. Leave packets use the existing signed wire format and are accepted only when the signature matches the key learned from a verified announcement. Invalid or