diff --git a/Tun/Context/SDLContextActor.swift b/Tun/Context/SDLContextActor.swift index 07b41dd..a1be1ee 100644 --- a/Tun/Context/SDLContextActor.swift +++ b/Tun/Context/SDLContextActor.swift @@ -9,6 +9,51 @@ import Foundation import NetworkExtension import NIOCore +private actor SDLWorkerRestartSignal { + private var generation: UInt64 = 0 + private var waiters: [UUID: CheckedContinuation] = [:] + + func request() { + self.generation &+= 1 + let generation = self.generation + let waiters = self.waiters + self.waiters.removeAll() + + for waiter in waiters.values { + waiter.resume(returning: generation) + } + } + + func currentGeneration() -> UInt64 { + return self.generation + } + + func waitForChange(after observedGeneration: UInt64) async -> UInt64 { + if self.generation != observedGeneration { + return self.generation + } + + let id = UUID() + return await withTaskCancellationHandler { + await withCheckedContinuation { continuation in + if self.generation != observedGeneration { + continuation.resume(returning: self.generation) + } else { + self.waiters[id] = continuation + } + } + } onCancel: { + Task { + await self.cancelWaiter(id: id) + } + } + } + + private func cancelWaiter(id: UUID) { + self.waiters.removeValue(forKey: id) + } +} + // 上下文环境变量,全局共享 /* 1. 处理rsa的加解密逻辑 @@ -65,6 +110,11 @@ actor SDLContextActor { private var rootTaskID: UUID? private var terminalError: Error? private let readySignal = AsyncOneShot() + private let superRestartSignal = SDLWorkerRestartSignal() + private let udpHoleRestartSignal = SDLWorkerRestartSignal() + private let udpHoleV6RestartSignal = SDLWorkerRestartSignal() + private let dnsCloudRestartSignal = SDLWorkerRestartSignal() + private let dnsLocalRestartSignal = SDLWorkerRestartSignal() public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) { let puncherActor = SDLPuncherActor() @@ -209,12 +259,15 @@ actor SDLContextActor { throw TunnelError.invalidContext } - let prunedSessions = await self.sessionManager.pruneExpiredSessions() + let clearedSessions = await self.sessionManager.clear() + self.natType = .blocked + await self.stopCurrentIPv6AssistClient() await self.packetOutboundActor.updateRuntime(config: self.config, dataCipher: dataCipher) await self.packetInboundActor.updateRuntime(config: self.config, dataCipher: dataCipher) try await self.tunNetworkManager.apply(settings: .init(config: self.config), dnsServer: DNSHelper.dnsServer) + await self.restartVolatileResourcesAfterWake() - SDLLogger.log("[SDLContext] recoverAfterWake completed, prunedSessions: \(prunedSessions)", category: .context) + SDLLogger.log("[SDLContext] recoverAfterWake completed, clearedSessions: \(clearedSessions)", category: .context) } private func runRootBody() async throws { @@ -266,6 +319,11 @@ actor SDLContextActor { let puncherActor = self.puncherActor let arpResolver = self.arpResolver let readySignal = self.readySignal + let superRestartSignal = self.superRestartSignal + let udpHoleRestartSignal = self.udpHoleRestartSignal + let udpHoleV6RestartSignal = self.udpHoleV6RestartSignal + let dnsCloudRestartSignal = self.dnsCloudRestartSignal + let dnsLocalRestartSignal = self.dnsLocalRestartSignal try await withThrowingTaskGroup(of: Void.self) { group in defer { @@ -273,31 +331,31 @@ actor SDLContextActor { } group.addTask { - try await Self.runRestarting(name: "superService") { + try await Self.runRestarting(name: "superService", restartSignal: superRestartSignal) { try await superService.run() } } group.addTask { - try await Self.runRestarting(name: "udpHoleService") { + try await Self.runRestarting(name: "udpHoleService", restartSignal: udpHoleRestartSignal) { try await udpHoleService.run() } } group.addTask { - try await Self.runRestarting(name: "udpHoleV6Service") { + try await Self.runRestarting(name: "udpHoleV6Service", restartSignal: udpHoleV6RestartSignal) { try await udpHoleV6Service.run() } } group.addTask { - try await Self.runRestarting(name: "dnsCloudService") { + try await Self.runRestarting(name: "dnsCloudService", restartSignal: dnsCloudRestartSignal) { try await dnsCloudService.run() } } group.addTask { - try await Self.runRestarting(name: "dnsLocalService") { + try await Self.runRestarting(name: "dnsLocalService", restartSignal: dnsLocalRestartSignal) { try await dnsLocalService.run() } } @@ -381,6 +439,22 @@ actor SDLContextActor { await client?.stop() } + private func restartVolatileResourcesAfterWake() async { + SDLLogger.log("[SDLContext] restart volatile resources after wake", category: .context) + + await self.superRestartSignal.request() + await self.udpHoleRestartSignal.request() + await self.udpHoleV6RestartSignal.request() + await self.dnsCloudRestartSignal.request() + await self.dnsLocalRestartSignal.request() + + await self.superService.stop() + await self.udpHoleService.stop() + await self.udpHoleV6Service.stop() + await self.dnsCloudService.stop() + await self.dnsLocalService.stop() + } + private func cleanupRoot() async { await self.puncherActor.stop() await self.arpResolver.stop() @@ -613,8 +687,11 @@ extension SDLContextActor { private static func runRestarting( name: String, retryDelay: Duration = .seconds(5), + restartSignal: SDLWorkerRestartSignal? = nil, operation: @escaping @Sendable () async throws -> Void ) async throws { + var restartGeneration = await restartSignal?.currentGeneration() ?? 0 + while !Task.isCancelled { do { try Task.checkCancellation() @@ -627,7 +704,45 @@ extension SDLContextActor { SDLLogger.log("[SDLContext] worker \(name) crashed: \(error.localizedDescription), will restart", category: .context) } + let nextGeneration = try await Self.waitForRestartSignalOrDelay( + name: name, + retryDelay: retryDelay, + restartSignal: restartSignal, + observedGeneration: restartGeneration + ) + restartGeneration = nextGeneration + } + } + + private static func waitForRestartSignalOrDelay( + name: String, + retryDelay: Duration, + restartSignal: SDLWorkerRestartSignal?, + observedGeneration: UInt64 + ) async throws -> UInt64 { + guard let restartSignal else { try await Task.sleep(for: retryDelay) + return observedGeneration + } + + return try await withThrowingTaskGroup(of: UInt64.self) { group in + group.addTask { + try await Task.sleep(for: retryDelay) + return await restartSignal.currentGeneration() + } + + group.addTask { + return await restartSignal.waitForChange(after: observedGeneration) + } + + let nextGeneration = try await group.next() ?? observedGeneration + group.cancelAll() + + if nextGeneration != observedGeneration { + SDLLogger.log("[SDLContext] worker \(name) received restart signal", category: .context) + } + + return nextGeneration } } diff --git a/Tun/Context/SDLContextBootstrap.swift b/Tun/Context/SDLContextBootstrap.swift index d1657ef..882bd59 100644 --- a/Tun/Context/SDLContextBootstrap.swift +++ b/Tun/Context/SDLContextBootstrap.swift @@ -40,7 +40,9 @@ final class SDLContextBootstrap: @unchecked Sendable { self.provider = provider self.commandContinuation = commandPair.continuation - self.commandWorker = Task { [weak self, stream = commandPair.stream] in + let stream = commandPair.stream + + self.commandWorker = Task { [weak self] in await self?.runCommandLoop(stream) } } diff --git a/Tun/Session/SessionTable.swift b/Tun/Session/SessionTable.swift index fb1c750..f8b6d91 100644 --- a/Tun/Session/SessionTable.swift +++ b/Tun/Session/SessionTable.swift @@ -66,9 +66,12 @@ actor SessionManager { self.publishSnapshot() } - func clear() { + @discardableResult + func clear() -> Int { + let oldCount = self.sessionCount() self.sessions.removeAll() self.publishSnapshot() + return oldCount } @discardableResult