diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/ServiceManager.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/ServiceManager.kt index 014dfa54..7e4095a0 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/ServiceManager.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/service/ServiceManager.kt @@ -3,29 +3,35 @@ package com.zaneschepke.wireguardautotunnel.core.service import android.app.Service import android.content.Context import android.content.Intent -import com.zaneschepke.wireguardautotunnel.domain.repository.AppDataRepository import com.zaneschepke.wireguardautotunnel.core.service.autotunnel.AutoTunnelService import com.zaneschepke.wireguardautotunnel.core.service.tile.AutoTunnelControlTile import com.zaneschepke.wireguardautotunnel.core.service.tile.TunnelControlTile +import com.zaneschepke.wireguardautotunnel.di.ApplicationScope +import com.zaneschepke.wireguardautotunnel.di.IoDispatcher import com.zaneschepke.wireguardautotunnel.domain.entity.TunnelConf +import com.zaneschepke.wireguardautotunnel.domain.repository.AppDataRepository import com.zaneschepke.wireguardautotunnel.util.extensions.requestAutoTunnelTileServiceUpdate import com.zaneschepke.wireguardautotunnel.util.extensions.requestTunnelTileServiceStateUpdate import jakarta.inject.Inject import kotlinx.coroutines.CompletableDeferred import kotlinx.coroutines.CoroutineDispatcher -import kotlinx.coroutines.ExperimentalCoroutinesApi +import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.asStateFlow import kotlinx.coroutines.flow.update +import kotlinx.coroutines.launch import kotlinx.coroutines.withContext +import kotlinx.coroutines.withTimeoutOrNull import timber.log.Timber -@OptIn(ExperimentalCoroutinesApi::class) -class ServiceManager -@Inject constructor(private val context: Context, private val ioDispatcher: CoroutineDispatcher, private val appDataRepository: AppDataRepository) { +class ServiceManager @Inject constructor( + private val context: Context, + @IoDispatcher private val ioDispatcher: CoroutineDispatcher, + @ApplicationScope private val applicationScope: CoroutineScope, + private val appDataRepository: AppDataRepository, +) { private val _autoTunnelActive = MutableStateFlow(false) - val autoTunnelActive = _autoTunnelActive.asStateFlow() var autoTunnelService = CompletableDeferred() @@ -44,76 +50,111 @@ class ServiceManager }.onFailure { Timber.e(it) } } - suspend fun startAutoTunnel(background: Boolean) { - val settings = appDataRepository.settings.get() - appDataRepository.settings.save(settings.copy(isAutoTunnelEnabled = true)) - if (autoTunnelService.isCompleted) return _autoTunnelActive.update { true } - runCatching { - startService(AutoTunnelService::class.java, background) - autoTunnelService.await() - autoTunnelService.getCompleted().start() - _autoTunnelActive.update { true } - updateAutoTunnelTile() - }.onFailure { - Timber.e(it) + fun startAutoTunnel(background: Boolean) { + applicationScope.launch(ioDispatcher) { + val settings = appDataRepository.settings.get() + appDataRepository.settings.save(settings.copy(isAutoTunnelEnabled = true)) + if (autoTunnelService.isCompleted && autoTunnelService.isActive) { + _autoTunnelActive.update { true } + return@launch + } + runCatching { + autoTunnelService = CompletableDeferred() // Reset + startService(AutoTunnelService::class.java, background) + val service = withTimeoutOrNull(SERVICE_START_TIMEOUT) { autoTunnelService.await() } + ?: throw IllegalStateException("AutoTunnelService start timed out") + service.start() + _autoTunnelActive.update { true } + updateAutoTunnelTile() + }.onFailure { + Timber.e(it) + _autoTunnelActive.update { false } + } } } - suspend fun startBackgroundService(tunnelConf: TunnelConf) { - if (backgroundService.isCompleted) return - runCatching { - startService(TunnelForegroundService::class.java, true) - backgroundService.await() - backgroundService.getCompleted().start(tunnelConf) - }.onFailure { - Timber.e(it) + fun startBackgroundService(tunnelConf: TunnelConf) { + applicationScope.launch(ioDispatcher) { + if (backgroundService.isCompleted && backgroundService.isActive) return@launch + runCatching { + backgroundService = CompletableDeferred() + startService(TunnelForegroundService::class.java, true) + val service = withTimeoutOrNull(SERVICE_START_TIMEOUT) { backgroundService.await() } + ?: throw IllegalStateException("Background service start timed out") + service.start(tunnelConf) + }.onFailure { + Timber.e(it) + } } } fun stopBackgroundService() { - if (!backgroundService.isCompleted) return - runCatching { - backgroundService.getCompleted().stop() - }.onFailure { - Timber.e(it) + applicationScope.launch(ioDispatcher) { + if (!backgroundService.isCompleted || !backgroundService.isActive) return@launch + runCatching { + val service = backgroundService.await() + service.stop() + backgroundService = CompletableDeferred() + }.onFailure { + Timber.e(it) + } } } - suspend fun toggleAutoTunnel(background: Boolean) { + fun toggleAutoTunnel(background: Boolean) { + applicationScope.launch(ioDispatcher) { + if (_autoTunnelActive.value) stopAutoTunnel() else startAutoTunnel(background) + } + } + + suspend fun updateAutoTunnelTile() { withContext(ioDispatcher) { - if (_autoTunnelActive.value) return@withContext stopAutoTunnel() - startAutoTunnel(background) + runCatching { + val service = withTimeoutOrNull(SERVICE_START_TIMEOUT) { autoTunnelTile.await() } + ?: run { + context.requestAutoTunnelTileServiceUpdate() + return@withContext + } + service.updateTileState() + }.onFailure { + Timber.e(it) + } } } - fun updateAutoTunnelTile() { - if (autoTunnelTile.isCompleted) { - autoTunnelTile.getCompleted().updateTileState() - } else { - context.requestAutoTunnelTileServiceUpdate() - } - } - - fun updateTunnelTile() { - if (tunnelControlTile.isCompleted) { - tunnelControlTile.getCompleted().updateTileState() - } else { - context.requestTunnelTileServiceStateUpdate() - } - } - - suspend fun stopAutoTunnel() { + suspend fun updateTunnelTile() { withContext(ioDispatcher) { + runCatching { + val service = withTimeoutOrNull(SERVICE_START_TIMEOUT) { tunnelControlTile.await() } + ?: run { + context.requestTunnelTileServiceStateUpdate() + return@withContext + } + service.updateTileState() + }.onFailure { + Timber.e(it) + } + } + } + + fun stopAutoTunnel() { + applicationScope.launch(ioDispatcher) { val settings = appDataRepository.settings.get() appDataRepository.settings.save(settings.copy(isAutoTunnelEnabled = false)) - if (!autoTunnelService.isCompleted) return@withContext + if (!autoTunnelService.isCompleted || !autoTunnelService.isActive) return@launch runCatching { - autoTunnelService.getCompleted().stop() + val service = autoTunnelService.await() + service.stop() _autoTunnelActive.update { false } + autoTunnelService = CompletableDeferred() updateAutoTunnelTile() }.onFailure { Timber.e(it) } } } + + companion object { + const val SERVICE_START_TIMEOUT = 5_000L + } } 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 354394ac..7bc2751e 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 @@ -177,16 +177,18 @@ open class BaseTunnel( private suspend fun startPingJob(tunnel: TunnelConf) = coroutineScope { while (isActive) { - if (isNetworkAvailable.get() && tunnel.isActive) { - val pingResult = tunnel.pingTunnel(ioDispatcher) - handlePingResult(tunnel, pingResult) + runCatching { + if (isNetworkAvailable.get() && tunnel.isActive) { + val pingSuccess = tunnel.isTunnelPingable(ioDispatcher) + handlePingResult(tunnel, pingSuccess) + } + delay(tunnel.pingInterval ?: Constants.PING_INTERVAL) } - delay(tunnel.pingInterval ?: Constants.PING_INTERVAL) } } - private suspend fun handlePingResult(tunnel: TunnelConf, pingResult: List) { - if (pingResult.contains(false)) { + private suspend fun handlePingResult(tunnel: TunnelConf, pingSuccess: Boolean) { + if (!pingSuccess) { if (isNetworkAvailable.get()) { Timber.i("Ping result: target was not reachable, bouncing the tunnel") bounceTunnel(tunnel) @@ -231,11 +233,13 @@ open class BaseTunnel( private suspend fun startTunnelStatisticsJob(tunnel: TunnelConf) = coroutineScope { while (isActive) { - val stats = getStatistics(tunnel) - tunnel.state.update { - it.copy(statistics = stats) + runCatching { + val stats = getStatistics(tunnel) + tunnel.state.update { + it.copy(statistics = stats) + } + delay(CHECK_INTERVAL) } - delay(CHECK_INTERVAL) } } } diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/KernelTunnel.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/KernelTunnel.kt index 5c9ac87b..62fa0641 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/KernelTunnel.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/KernelTunnel.kt @@ -32,12 +32,16 @@ class KernelTunnel @Inject constructor( ) : BaseTunnel(ioDispatcher, applicationScope, networkMonitor, appDataRepository, serviceManager, notificationManager) { override fun startTunnel(tunnelConf: TunnelConf) { + Timber.d("Starting tunnel ${tunnelConf.id} kernel") applicationScope.launch(ioDispatcher) { if (tunnels.value.any { it.id == tunnelConf.id }) return@launch Timber.w("Tunnel already running") runCatching { + Timber.d("Setting backend state UP") backend.setState(tunnelConf, Tunnel.State.UP, tunnelConf.toWgConfig()) + Timber.d("Calling super.startTunnel") super.startTunnel(tunnelConf) }.onFailure { + Timber.e(it, "Failed to start tunnel ${tunnelConf.id} kernel") onTunnelStop(tunnelConf) if (it is BackendException) { handleBackendThrowable(it.toBackendError()) @@ -78,6 +82,13 @@ class KernelTunnel @Inject constructor( } } + override suspend fun bounceTunnel(tunnelConf: TunnelConf) { + if (tunnels.value.any { it.id == tunnelConf.id }) { + toggleTunnel(tunnelConf, TunnelStatus.DOWN) + toggleTunnel(tunnelConf, TunnelStatus.UP) + } + } + override suspend fun setBackendState(backendState: BackendState, allowedIps: Collection) { Timber.w("Not yet implemented for kernel") } diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/UserspaceTunnel.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/UserspaceTunnel.kt index 2caa56bb..00c52e90 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/UserspaceTunnel.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/core/tunnel/UserspaceTunnel.kt @@ -34,14 +34,19 @@ class UserspaceTunnel @Inject constructor( override fun startTunnel(tunnelConf: TunnelConf) { applicationScope.launch(ioDispatcher) { + Timber.d("Starting tunnel ${tunnelConf.id} userspace") if (tunnels.value.any { it.id == tunnelConf.id }) return@launch Timber.w("Tunnel already running") if (tunnels.value.isNotEmpty()) { + Timber.d("Stopping all tunnels") stopAllTunnels() } runCatching { + Timber.d("Setting backend state UP") backend.setState(tunnelConf, Tunnel.State.UP, tunnelConf.toAmConfig()) + Timber.d("Calling super.startTunnel") super.startTunnel(tunnelConf) }.onFailure { + Timber.e(it, "Failed to start tunnel ${tunnelConf.id} userspace") onTunnelStop(tunnelConf) if (it is BackendException) { handleBackendThrowable(it.toBackendError()) diff --git a/app/src/main/java/com/zaneschepke/wireguardautotunnel/domain/entity/TunnelConf.kt b/app/src/main/java/com/zaneschepke/wireguardautotunnel/domain/entity/TunnelConf.kt index 32519a6c..fbc92535 100644 --- a/app/src/main/java/com/zaneschepke/wireguardautotunnel/domain/entity/TunnelConf.kt +++ b/app/src/main/java/com/zaneschepke/wireguardautotunnel/domain/entity/TunnelConf.kt @@ -74,18 +74,17 @@ data class TunnelConf( updatedConf.pingInterval == pingInterval } - suspend fun pingTunnel(context: CoroutineContext): List { + suspend fun isTunnelPingable(context: CoroutineContext): Boolean { return withContext(context) { val config = toWgConfig() if (pingIp != null) { - Timber.i("Pinging custom ip") - listOf(InetAddress.getByName(pingIp).isReachable(Constants.PING_TIMEOUT.toInt())) - } else { - Timber.i("Pinging all peers") - config.peers.map { peer -> - peer.isReachable(isIpv4Preferred) - } + return@withContext InetAddress.getByName(pingIp) + .isReachable(Constants.PING_TIMEOUT.toInt()) } + Timber.i("Pinging all peers") + config.peers.map { peer -> + peer.isReachable(isIpv4Preferred) + }.all { true } } }