From 5ff23477562bfab9bd49e4a28eec8fd5b81bbbd5 Mon Sep 17 00:00:00 2001 From: hawkff <109485367+hawkff@users.noreply.github.com> Date: Sun, 12 Jul 2026 17:56:08 -0400 Subject: [PATCH] fix(network): preserve listener after fallback --- .../sagernet/utils/DefaultNetworkListener.kt | 40 ++++++- .../utils/NetworkCallbackRegistration.kt | 33 ++++++ .../utils/NetworkCallbackRegistrationTest.kt | 107 ++++++++++++++++++ 3 files changed, 174 insertions(+), 6 deletions(-) create mode 100644 app/src/main/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistration.kt create mode 100644 app/src/test/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistrationTest.kt diff --git a/app/src/main/java/io/nekohasekai/sagernet/utils/DefaultNetworkListener.kt b/app/src/main/java/io/nekohasekai/sagernet/utils/DefaultNetworkListener.kt index 9c3eb9c081..767d6e184f 100644 --- a/app/src/main/java/io/nekohasekai/sagernet/utils/DefaultNetworkListener.kt +++ b/app/src/main/java/io/nekohasekai/sagernet/utils/DefaultNetworkListener.kt @@ -42,7 +42,16 @@ object DefaultNetworkListener { is NetworkMessage.Start -> { if (listeners.isEmpty()) register() listeners[message.key] = message.listener - if (network != null) message.listener(network) + if (network != null) { + message.listener(network) + } else if (fallback) { + val activeNetwork = if (Build.VERSION.SDK_INT >= Build.VERSION_CODES.M) { + SagerNet.connectivity.activeNetwork + } else { + null + } + message.listener(activeNetwork) + } } is NetworkMessage.Get -> { check(listeners.isNotEmpty()) { "Getting network without any listeners is not supported" } @@ -124,6 +133,7 @@ object DefaultNetworkListener { } private var fallback = false + private val callbackRegistration = NetworkCallbackRegistration() private val request = NetworkRequest.Builder().apply { addCapability(NetworkCapabilities.NET_CAPABILITY_INTERNET) addCapability(NetworkCapabilities.NET_CAPABILITY_NOT_RESTRICTED) @@ -145,8 +155,19 @@ object DefaultNetworkListener { * Source: https://android.googlesource.com/platform/frameworks/base/+/2df4c7d/services/core/java/com/android/server/ConnectivityService.java#887 */ private fun register() { - try { - fallback = false + if (callbackRegistration.requiresFallback) { + callbackRegistration.unregister { + SagerNet.connectivity.unregisterNetworkCallback(Callback) + }.onFailure { + Logs.w("DefaultNetworkListener: retry unregister failed", it) + } + } + if (callbackRegistration.requiresFallback) { + fallback = true + return + } + fallback = false + callbackRegistration.register { when (Build.VERSION.SDK_INT) { in 31..Int.MAX_VALUE -> @TargetApi(31) @@ -177,11 +198,18 @@ object DefaultNetworkListener { // known bug on API 23: https://stackoverflow.com/a/33509180/2245107 } } - } catch (e: Exception) { - Logs.w(e) + }.onFailure { + Logs.w(it) fallback = true } } - private fun unregister() = SagerNet.connectivity.unregisterNetworkCallback(Callback) + private fun unregister() { + callbackRegistration.unregister { + SagerNet.connectivity.unregisterNetworkCallback(Callback) + }.onFailure { + fallback = true + Logs.w("DefaultNetworkListener: failed to unregister network callback", it) + } + } } diff --git a/app/src/main/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistration.kt b/app/src/main/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistration.kt new file mode 100644 index 0000000000..3db3edf736 --- /dev/null +++ b/app/src/main/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistration.kt @@ -0,0 +1,33 @@ +package io.nekohasekai.sagernet.utils + +internal class NetworkCallbackRegistration { + internal var isRegistered = false + private set + internal var requiresFallback = false + private set + + fun register(block: () -> Unit): Result { + if (isRegistered) return Result.success(Unit) + return runCatching(block).onSuccess { + isRegistered = true + requiresFallback = false + } + } + + fun unregister(block: () -> Unit): Result { + if (!isRegistered) return Result.success(Unit) + return runCatching(block) + .onSuccess { + isRegistered = false + requiresFallback = false + } + .onFailure { throwable -> + if (throwable is IllegalArgumentException) { + isRegistered = false + requiresFallback = false + } else { + requiresFallback = true + } + } + } +} diff --git a/app/src/test/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistrationTest.kt b/app/src/test/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistrationTest.kt new file mode 100644 index 0000000000..9adf28558d --- /dev/null +++ b/app/src/test/java/io/nekohasekai/sagernet/utils/NetworkCallbackRegistrationTest.kt @@ -0,0 +1,107 @@ +package io.nekohasekai.sagernet.utils + +import org.junit.Assert.assertEquals +import org.junit.Assert.assertFalse +import org.junit.Assert.assertSame +import org.junit.Assert.assertTrue +import org.junit.Test + +class NetworkCallbackRegistrationTest { + + @Test + fun successfulRegisterThenUnregisterInvokesEachLambdaOnce() { + val registration = NetworkCallbackRegistration() + var registerCalls = 0 + var unregisterCalls = 0 + + val registerResult = registration.register { registerCalls++ } + val unregisterResult = registration.unregister { unregisterCalls++ } + + assertTrue(registerResult.isSuccess) + assertTrue(unregisterResult.isSuccess) + assertEquals(1, registerCalls) + assertEquals(1, unregisterCalls) + assertFalse(registration.isRegistered) + } + + @Test + fun failedRegisterLeavesStateFalseAndSkipsUnregisterLambda() { + val registration = NetworkCallbackRegistration() + val failure = IllegalStateException("registration failed") + var unregisterCalls = 0 + + val registerResult = registration.register { throw failure } + val unregisterResult = registration.unregister { unregisterCalls++ } + + assertSame(failure, registerResult.exceptionOrNull()) + assertTrue(unregisterResult.isSuccess) + assertEquals(0, unregisterCalls) + assertFalse(registration.isRegistered) + } + + @Test + fun failedUnregisterKeepsStateForRetry() { + val registration = NetworkCallbackRegistration() + val failure = IllegalStateException("unregistration failed") + var successfulUnregisterCalls = 0 + registration.register {} + + val failedResult = registration.unregister { throw failure } + val retryResult = registration.unregister { successfulUnregisterCalls++ } + + assertSame(failure, failedResult.exceptionOrNull()) + assertTrue(retryResult.isSuccess) + assertEquals(1, successfulUnregisterCalls) + assertFalse(registration.isRegistered) + assertFalse(registration.requiresFallback) + } + + @Test + fun alreadyUnregisteredFailureAllowsFreshRegistration() { + val registration = NetworkCallbackRegistration() + var registerCalls = 0 + registration.register {} + + val unregisterResult = registration.unregister { + throw IllegalArgumentException("callback was not registered") + } + val registerResult = registration.register { registerCalls++ } + + assertTrue(unregisterResult.isFailure) + assertTrue(registerResult.isSuccess) + assertEquals(1, registerCalls) + assertTrue(registration.isRegistered) + assertFalse(registration.requiresFallback) + } + + @Test + fun registerAfterEitherFailureDoesNotDuplicateARegisteredCallback() { + val registration = NetworkCallbackRegistration() + var successfulRegisterCalls = 0 + + registration.register { throw IllegalStateException("registration failed") } + val afterRegisterFailure = registration.register { successfulRegisterCalls++ } + registration.unregister { throw IllegalStateException("unregistration failed") } + val afterUnregisterFailure = registration.register { successfulRegisterCalls++ } + + assertTrue(afterRegisterFailure.isSuccess) + assertTrue(afterUnregisterFailure.isSuccess) + assertEquals(1, successfulRegisterCalls) + assertTrue(registration.isRegistered) + assertTrue(registration.requiresFallback) + } + + @Test + fun repeatedUnregisterWhileFalseIsHarmless() { + val registration = NetworkCallbackRegistration() + var unregisterCalls = 0 + + val firstResult = registration.unregister { unregisterCalls++ } + val secondResult = registration.unregister { unregisterCalls++ } + + assertTrue(firstResult.isSuccess) + assertTrue(secondResult.isSuccess) + assertEquals(0, unregisterCalls) + assertFalse(registration.isRegistered) + } +}