From 5ded556647bb024332a537cfa74e7f77f42850a2 Mon Sep 17 00:00:00 2001 From: Zane Schepke Date: Fri, 11 Apr 2025 17:54:40 -0400 Subject: [PATCH] fix: ping job start by default fix: kernel dns resolution on stop bug closes #674 --- .../core/service/TunnelForegroundService.kt | 156 ++++++++++++------ .../core/tunnel/BaseTunnel.kt | 67 +++++--- .../core/tunnel/TunnelManager.kt | 8 +- .../core/tunnel/TunnelProvider.kt | 18 ++ gradle/libs.versions.toml | 2 +- 5 files changed, 172 insertions(+), 79 deletions(-) diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/TunnelForegroundService.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/TunnelForegroundService.kt index 05950575..5289b40d 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/TunnelForegroundService.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/TunnelForegroundService.kt @@ -17,6 +17,7 @@ import com.zaneschepke.wireguardautotunnel.domain.entity.TunnelConf import com.zaneschepke.wireguardautotunnel.domain.enums.NotificationAction import com.zaneschepke.wireguardautotunnel.domain.enums.TunnelStatus import com.zaneschepke.wireguardautotunnel.domain.repository.TunnelRepository +import com.zaneschepke.wireguardautotunnel.domain.state.TunnelState import com.zaneschepke.wireguardautotunnel.util.Constants import com.zaneschepke.wireguardautotunnel.util.extensions.distinctByKeys import dagger.hilt.android.AndroidEntryPoint @@ -36,6 +37,8 @@ import kotlinx.coroutines.flow.flowOn import kotlinx.coroutines.flow.map import kotlinx.coroutines.isActive import kotlinx.coroutines.launch +import kotlinx.coroutines.sync.Mutex +import kotlinx.coroutines.sync.withLock import kotlinx.coroutines.withContext import timber.log.Timber @@ -57,6 +60,9 @@ class TunnelForegroundService : LifecycleService() { private val isNetworkConnected = MutableStateFlow(true) private val tunnelJobs = ConcurrentHashMap() + private val pingJobs = ConcurrentHashMap() + + private val jobsMutex = Mutex() override fun onCreate() { super.onCreate() @@ -88,31 +94,85 @@ class TunnelForegroundService : LifecycleService() { } fun start() = - lifecycleScope.launch { - tunnelManager.activeTunnels.distinctByKeys().collect { tuns -> - if (tuns.isEmpty() && tunnelJobs.isEmpty()) return@collect - if (tuns.isEmpty() && tunnelJobs.isNotEmpty()) { - return@collect tunnelJobs.forEach { (key, _) -> - Timber.d("Stopping all tunnel jobs") - tunnelJobs[key]?.cancel() - tunnelJobs.remove(key) - } - } - val (jobsToStop, jobsToStart) = findMissingKeys(tuns, tunnelJobs) - if (jobsToStop.isEmpty() && jobsToStart.isEmpty()) return@collect - jobsToStop.forEach { tun -> - Timber.d("Stopping tunnel jobs for ${tun.tunName}") - tunnelJobs[tun]?.cancel() - tunnelJobs.remove(tun) - } - jobsToStart.forEach { tun -> - Timber.d("Starting tunnel jobs for ${tun.tunName}") - tunnelJobs += (tun to startTunnelJobs(tun)) - } + lifecycleScope.launch(ioDispatcher) { + tunnelManager.activeTunnels.distinctByKeys().collect { activeTunnels -> + // No active tunnels and no jobs: nothing to do + if (activeTunnels.isEmpty() && tunnelJobs.isEmpty()) return@collect + + // Synchronize jobs with active tunnels + synchronizeJobs(activeTunnels) updateServiceNotification() } } + private suspend fun synchronizeJobs(activeTunnels: Map) { + jobsMutex.withLock { + // Stop jobs for tunnels that are no longer active + stopInactiveJobs(activeTunnels) + // Start jobs for new tunnels + startNewJobs(activeTunnels) + } + } + + private fun stopInactiveJobs(activeTunnels: Map) { + // If no active tunnels, clear all jobs + if (activeTunnels.isEmpty()) { + clearAllJobs() + return + } + // Stop jobs for tunnels not in activeTunnels + val tunnelsToStop = tunnelJobs.keys - activeTunnels.keys + tunnelsToStop.forEach { tun -> stopTunnelJobs(tun) } + } + + private fun clearAllJobs() { + tunnelJobs.forEach { (tun, job) -> + Timber.d("Stopping tunnel job for ${tun.tunName}") + job.cancel() + } + tunnelJobs.clear() + + pingJobs.forEach { (tun, job) -> + if (isPingBounce(tun)) { + Timber.d("Preserving ping job for ${tun.tunName} due to PING bounce") + return@forEach + } + Timber.d("Stopping ping job for ${tun.tunName}") + job.cancel() + } + pingJobs.entries.removeIf { (tun, _) -> !isPingBounce(tun) } + } + + private fun stopTunnelJobs(tun: TunnelConf) { + tunnelJobs.remove(tun)?.cancel() + Timber.d("Stopped tunnel job for ${tun.tunName}") + if (isPingBounce(tun)) + return Timber.d("Preserving ${tun.tunName} ping job due to ping bounce") + pingJobs.remove(tun)?.cancel() + Timber.d("Stopped ping job for ${tun.tunName}") + } + + private fun startNewJobs(activeTunnels: Map) { + val tunnelsToStart = activeTunnels.keys - tunnelJobs.keys + tunnelsToStart.forEach { tun -> + tunnelJobs[tun] = startTunnelJobs(tun) + Timber.d("Started tunnel job for ${tun.tunName}") + + if (pingJobs[tun]?.isActive == true) { + Timber.d("Reusing active ping job for ${tun.tunName}") + } else { + pingJobs[tun]?.cancel() // Cancel any stale job + if (tun.isPingEnabled) { + pingJobs[tun] = startPingJob(tun) + Timber.d("Started ping job for ${tun.tunName}") + } + } + } + } + + private fun isPingBounce(tun: TunnelConf): Boolean = + tunnelManager.bouncingTunnelIds[tun.id] == TunnelStatus.StopReason.PING + // TODO Would be cool to have this include kill switch // TODO also we need to include errors private fun updateServiceNotification() { @@ -132,26 +192,15 @@ class TunnelForegroundService : LifecycleService() { // use same scope so we can cancel all of these private fun startTunnelJobs(tunnelConf: TunnelConf) = - lifecycleScope.launch { + lifecycleScope.launch(ioDispatcher) { // monitor if we have internet connectivity launch { startNetworkMonitorJob() } // job to trigger stats emit on interval launch { startTunnelStatsJob(tunnelConf) } // monitor changes to the tunnel config launch { startTunnelConfChangesJob(tunnelConf) } - // monitor tunnel ping - launch { startPingJob(tunnelConf) } } - private fun findMissingKeys( - map1: Map, - map2: Map, - ): Pair, Set> { - val missingMap1 = map2.keys - map1.keys - val missingMap2 = map1.keys - map2.keys - return missingMap1 to missingMap2 - } - private suspend fun startTunnelConfChangesJob(tunnelConf: TunnelConf) { tunnelRepo.flow .flowOn(ioDispatcher) @@ -188,24 +237,25 @@ class TunnelForegroundService : LifecycleService() { } } - // TODO fix cooldown - private suspend fun startPingJob(tunnel: TunnelConf) = coroutineScope { - delay(PING_START_DELAY) - while (isActive) { - val shouldBounce = shouldBounceTunnel(tunnel) - val delayMs = - if (shouldBounce) { - // let this complete, even after cancel - withContext(NonCancellable) { - tunnelManager.bounceTunnel(tunnel, TunnelStatus.StopReason.PING) + private fun startPingJob(tunnel: TunnelConf) = + lifecycleScope.launch(ioDispatcher) { + // delay for initial duration + delay(tunnel.pingInterval ?: Constants.PING_INTERVAL) + while (isActive) { + val shouldBounce = shouldBounceTunnel(tunnel) + val delayMs = + if (shouldBounce) { + // let this complete, even after cancel + withContext(NonCancellable) { + tunnelManager.bounceTunnel(tunnel, TunnelStatus.StopReason.PING) + } + tunnel.pingCooldown ?: Constants.PING_COOLDOWN + } else { + tunnel.pingInterval ?: Constants.PING_INTERVAL } - tunnel.pingCooldown ?: Constants.PING_COOLDOWN - } else { - tunnel.pingInterval ?: Constants.PING_INTERVAL - } - delay(delayMs) + delay(delayMs) + } } - } private suspend fun shouldBounceTunnel(tunnel: TunnelConf): Boolean { if (!isNetworkConnected.value) { @@ -263,10 +313,10 @@ class TunnelForegroundService : LifecycleService() { // TODO add notification handling and optional log reading for restart on handshake failures companion object { const val STATS_DELAY = 1_000L - const val PING_START_DELAY = 30_000L // ipv6 disabled or block on network - // const val userspaceStartFailed = "Failed to send handshake initiation: write udp [::]" - // const val ipv6Fails = "Failed to send data packets: write udp [::]" - // const val ipv4Fails = "Failed to send data packets: write udp 0.0.0.0:51820" + // Failed to send handshake initiation: write udp [::]" + // Failed to send data packets: write udp [::] + // Failed to send data packets: write udp 0.0.0.0:51820 + // Handshake did not complete after 5 seconds, retrying } } diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/BaseTunnel.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/BaseTunnel.kt index 95337b04..e49e1f9b 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/BaseTunnel.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/BaseTunnel.kt @@ -1,5 +1,6 @@ package com.zaneschepke.wireguardautotunnel.core.tunnel +import com.wireguard.android.backend.BackendException import com.wireguard.android.backend.Tunnel import com.zaneschepke.wireguardautotunnel.core.service.ServiceManager import com.zaneschepke.wireguardautotunnel.di.ApplicationScope @@ -10,10 +11,11 @@ import com.zaneschepke.wireguardautotunnel.domain.repository.AppDataRepository import com.zaneschepke.wireguardautotunnel.domain.state.TunnelState import com.zaneschepke.wireguardautotunnel.domain.state.TunnelStatistics import com.zaneschepke.wireguardautotunnel.util.extensions.asTunnelState +import com.zaneschepke.wireguardautotunnel.util.extensions.toBackendError import java.util.concurrent.ConcurrentHashMap -import java.util.concurrent.atomic.AtomicBoolean import kotlin.concurrent.thread import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.delay import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.asStateFlow import kotlinx.coroutines.flow.update @@ -35,8 +37,9 @@ abstract class BaseTunnel( private val tunMutex = Mutex() private val tunStatusMutex = Mutex() + private val bounceTunnelMutex = Mutex() - private val isBouncing = AtomicBoolean(false) + override val bouncingTunnelIds = ConcurrentHashMap() abstract suspend fun startBackend(tunnel: TunnelConf) @@ -98,7 +101,7 @@ abstract class BaseTunnel( is org.amnezia.awg.backend.Tunnel.State -> updateTunnelStatus(tunnelConf, state.asTunnelState()) } - handleServiceChangesOnStop() + handleServiceStateOnChange() } serviceManager.updateTunnelTile() } @@ -111,13 +114,10 @@ abstract class BaseTunnel( override suspend fun startTunnel(tunnelConf: TunnelConf) { if (activeTuns.exists(tunnelConf.id) || tunThreads.containsKey(tunnelConf.id)) return - // stop active tunnels if we are userspace if (this@BaseTunnel is UserspaceTunnel) stopActiveTunnels() tunMutex.withLock { - // use thread to interrupt java backend if stuck (like in dns resolution) - tunThreads += - tunnelConf.id to - thread { + tunThreads[tunnelConf.id] = thread { + runCatching { runBlocking { try { Timber.d("Starting tunnel ${tunnelConf.id}...") @@ -127,23 +127,32 @@ abstract class BaseTunnel( Timber.e(e, "Failed to start tunnel ${tunnelConf.name} userspace") updateTunnelStatus(tunnelConf, TunnelStatus.Error(e)) } catch (e: InterruptedException) { - Timber.i( + Timber.w( "Tunnel start has been interrupted as ${tunnelConf.name} failed to start" ) } } } + .onFailure { Timber.w("Tunnel start has been interrupted") } + } } } private suspend fun startTunnelInner(tunnelConf: TunnelConf) { configureTunnelCallbacks(tunnelConf) - Timber.d("Started backend for tunnel ${tunnelConf.id}...") - startBackend(tunnelConf) - updateTunnelStatus(tunnelConf, TunnelStatus.Up) - Timber.d("DONE for tun ${tunnelConf.id}...") - saveTunnelActiveState(tunnelConf, true) - serviceManager.startTunnelForegroundService() + Timber.d("Starting backend for tunnel ${tunnelConf.id}...") + try { + startBackend(tunnelConf) + updateTunnelStatus(tunnelConf, TunnelStatus.Up) + Timber.d("Started for tun ${tunnelConf.id}...") + saveTunnelActiveState(tunnelConf, true) + serviceManager.startTunnelForegroundService() + } catch (e: BackendException) { + Timber.e(e, "Failed to start backend for ${tunnelConf.name}") + val backendError = e.toBackendError() + updateTunnelStatus(tunnelConf, TunnelStatus.Error(backendError)) + throw backendError + } } private suspend fun saveTunnelActiveState(tunnelConf: TunnelConf, active: Boolean) { @@ -173,9 +182,9 @@ abstract class BaseTunnel( removeActiveTunnel(tunnel) } - private suspend fun handleServiceChangesOnStop() { - if (activeTuns.value.isEmpty() && !isBouncing.get()) - return serviceManager.stopTunnelForegroundService() + private suspend fun handleServiceStateOnChange() { + if (activeTuns.value.isEmpty() && bouncingTunnelIds.isEmpty()) + serviceManager.stopTunnelForegroundService() } private suspend fun handleStuckStartingTunnelShutdown(tunnel: TunnelConf) { @@ -205,11 +214,23 @@ abstract class BaseTunnel( } override suspend fun bounceTunnel(tunnelConf: TunnelConf, reason: TunnelStatus.StopReason) { - Timber.i("Bounce tunnel ${tunnelConf.name}") - isBouncing.set(true) - stopTunnel(tunnelConf, reason) - startTunnel(tunnelConf) - isBouncing.set(false) + bounceTunnelMutex.withLock { + Timber.i( + "Bounce tunnel ${tunnelConf.name} for reason: $reason, current bouncing: ${bouncingTunnelIds.size}" + ) + bouncingTunnelIds[tunnelConf.id] = reason + try { + stopTunnel(tunnelConf, reason) + delay(300L) + startTunnel(tunnelConf) + } finally { + bouncingTunnelIds.remove(tunnelConf.id) + handleServiceStateOnChange() + Timber.d( + "Cleared bounce state for ${tunnelConf.name}, remaining: ${bouncingTunnelIds.size}" + ) + } + } } override suspend fun runningTunnelNames(): Set = diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelManager.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelManager.kt index 28795d5b..27da08bd 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelManager.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelManager.kt @@ -9,6 +9,7 @@ import com.zaneschepke.wireguardautotunnel.domain.enums.BackendState import com.zaneschepke.wireguardautotunnel.domain.enums.TunnelStatus import com.zaneschepke.wireguardautotunnel.domain.repository.AppDataRepository import com.zaneschepke.wireguardautotunnel.domain.state.TunnelStatistics +import java.util.concurrent.ConcurrentHashMap import javax.inject.Inject import kotlinx.coroutines.CoroutineDispatcher import kotlinx.coroutines.CoroutineScope @@ -61,6 +62,9 @@ constructor( initialValue = emptyMap(), ) + override val bouncingTunnelIds: ConcurrentHashMap = + tunnelProviderFlow.value.bouncingTunnelIds + override fun hasVpnPermission(): Boolean { return userspaceTunnel.hasVpnPermission() } @@ -78,11 +82,11 @@ constructor( } override suspend fun stopTunnel(tunnelConf: TunnelConf?, reason: TunnelStatus.StopReason) { - tunnelProviderFlow.value.stopTunnel(tunnelConf) + tunnelProviderFlow.value.stopTunnel(tunnelConf, reason) } override suspend fun bounceTunnel(tunnelConf: TunnelConf, reason: TunnelStatus.StopReason) { - tunnelProviderFlow.value.bounceTunnel(tunnelConf) + tunnelProviderFlow.value.bounceTunnel(tunnelConf, reason) } override fun setBackendState(backendState: BackendState, allowedIps: Collection) { diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelProvider.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelProvider.kt index e2600dc6..b8351691 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelProvider.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/TunnelProvider.kt @@ -5,16 +5,32 @@ import com.zaneschepke.wireguardautotunnel.domain.enums.BackendState import com.zaneschepke.wireguardautotunnel.domain.enums.TunnelStatus import com.zaneschepke.wireguardautotunnel.domain.state.TunnelState import com.zaneschepke.wireguardautotunnel.domain.state.TunnelStatistics +import java.util.concurrent.ConcurrentHashMap import kotlinx.coroutines.flow.StateFlow interface TunnelProvider { + /** Starts the specified tunnel configuration. */ suspend fun startTunnel(tunnelConf: TunnelConf) + /** + * Stops the specified tunnel, or all tunnels if none is provided. + * + * @param tunnelConf The tunnel to stop, or null to stop all active tunnels. + * @param reason The reason for stopping, defaults to USER for manual stops. Callers should + * override with specific reasons (e.g., PING, CONFIG_CHANGED) when applicable. + */ suspend fun stopTunnel( tunnelConf: TunnelConf? = null, reason: TunnelStatus.StopReason = TunnelStatus.StopReason.USER, ) + /** + * Bounces (stops and restarts) the specified tunnel. + * + * @param tunnelConf The tunnel to bounce. + * @param reason The reason for bouncing, defaults to USER for manual actions. Callers should + * override with specific reasons (e.g., PING, CONFIG_CHANGED) when applicable. + */ suspend fun bounceTunnel( tunnelConf: TunnelConf, reason: TunnelStatus.StopReason = TunnelStatus.StopReason.USER, @@ -30,6 +46,8 @@ interface TunnelProvider { val activeTunnels: StateFlow> + val bouncingTunnelIds: ConcurrentHashMap + fun hasVpnPermission(): Boolean suspend fun clearError(tunnelConf: TunnelConf) diff --git a/gradle/libs.versions.toml b/gradle/libs.versions.toml index 6e0c0cc6..b127f73b 100644 --- a/gradle/libs.versions.toml +++ b/gradle/libs.versions.toml @@ -19,7 +19,7 @@ navigationCompose = "2.8.9" pinLockCompose = "1.0.4" roomVersion = "2.7.0" timber = "5.0.1" -tunnel = "1.2.10" +tunnel = "1.2.11" androidGradlePlugin = "8.8.0-alpha05" kotlin = "2.1.20" ksp = "2.1.20-2.0.0"