diff --git a/Tun/Arp/ArpResolver.swift b/Tun/Arp/ArpResolver.swift index 196698b..aead8f1 100644 --- a/Tun/Arp/ArpResolver.swift +++ b/Tun/Arp/ArpResolver.swift @@ -20,24 +20,17 @@ actor ArpResolver { private var known_macs: [UInt32: ArpEntry] = [:] private let arpTTL: TimeInterval nonisolated private let snapshotPublisher: SnapshotPublisher - - private var cleanupTask: Task? - + init(arpTTL: TimeInterval = 300) { self.arpTTL = arpTTL self.snapshotPublisher = SnapshotPublisher(initial: ArpSnapshot.empty()) } - func start() { - guard self.cleanupTask == nil else { - return - } - - self.cleanupTask = Task { [weak self] in - while !Task.isCancelled { - try? await Task.sleep(for: .seconds(1)) - await self?.cleanup() - } + func runCleanup() async throws { + while !Task.isCancelled { + try await Task.sleep(for: .seconds(1)) + try Task.checkCancellation() + self.cleanup() } } @@ -78,8 +71,6 @@ actor ArpResolver { } func stop() { - self.cleanupTask?.cancel() - self.cleanupTask = nil self.clear() } @@ -132,8 +123,4 @@ actor ArpResolver { return ArpSnapshot(entries: entries) } - deinit { - self.cleanupTask?.cancel() - } - } diff --git a/Tun/Context/SDLContextActor.swift b/Tun/Context/SDLContextActor.swift index 02f8599..ce32580 100644 --- a/Tun/Context/SDLContextActor.swift +++ b/Tun/Context/SDLContextActor.swift @@ -14,29 +14,6 @@ import NIOCore 1. 处理rsa的加解密逻辑 */ -func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task { - return Task(name: name) { - while true { - do { - try Task.checkCancellation() - try await body() - } catch is CancellationError { - SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug) - break - } catch let err { - SDLLogger.log("[SDLContext] worker \(name) crashed: \(err.localizedDescription), will restart", for: .debug) - do { - try await Task.sleep(for: retryDelay) - } catch is CancellationError { - break - } catch { - break - } - } - } - } -} - enum SDLContextError: Error { case udpHoleClosed @@ -92,10 +69,6 @@ actor SDLContextActor { // 处理权限控制 private let policyService: PolicyService - private var updatePolicyWorker: PeriodicWorker? - - // stunRequest任务 - private var stunRequestWorker: PeriodicWorker? private var rootTask: Task? private let readySignal = AsyncOneShot() @@ -201,16 +174,11 @@ actor SDLContextActor { private func runRootBody() async throws { self.prepareTunnelNotifier() - // 启动arp的定时清理任务 - await self.puncherActor.start() - await self.arpResolver.start() - let dnsService = DNSService(serverIP: self.config.serverEndpoint.ip, publicDnsServers: self.publicDnsServers) { [weak self] event in await self?.handleDNSEvent(event) } self.dnsService = dnsService await self.packetOutboundActor.updateDNSService(dnsService) - await dnsService.start() await self.superService.updateMessageHandler { [weak self] message in await self?.handleSuperMessage(message: message) @@ -229,6 +197,9 @@ actor SDLContextActor { let superService = self.superService let udpHoleService = self.udpHoleService let packetOutboundActor = self.packetOutboundActor + let policyService = self.policyService + let puncherActor = self.puncherActor + let arpResolver = self.arpResolver let readySignal = self.readySignal try await withThrowingTaskGroup(of: Void.self) { group in @@ -248,11 +219,38 @@ actor SDLContextActor { } } + group.addTask { + try await dnsService.run() + } + + group.addTask { + try await puncherActor.runCleanup() + } + + group.addTask { + try await arpResolver.runCleanup() + } + group.addTask(priority: .high) { _ = try await readySignal.wait() try await packetOutboundActor.runPacketReader() } + group.addTask { + _ = try await readySignal.wait() + try await Self.runPeriodic(name: "updatePolicyTask", interval: .seconds(10)) { + SDLLogger.log("[SDLContext] updatePolicyTask execute") + await policyService.updatePolicy(superService: superService) + } + } + + group.addTask { + _ = try await readySignal.wait() + try await Self.runPeriodic(name: "stunRequestTask", interval: .seconds(8)) { + try await self.runStunRequestOnce() + } + } + try await group.waitForAll() } } @@ -278,16 +276,32 @@ actor SDLContextActor { } } + private static func runPeriodic( + name: String, + interval: Duration, + retryDelay: Duration = .seconds(5), + operation: @escaping @Sendable () async throws -> Void + ) async throws { + while !Task.isCancelled { + do { + try Task.checkCancellation() + try await operation() + try await Task.sleep(for: interval) + } catch is CancellationError { + SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug) + throw CancellationError() + } catch { + SDLLogger.log("[SDLContext] worker \(name) crashed: \(error.localizedDescription), will retry", for: .debug) + try await Task.sleep(for: retryDelay) + } + } + } + private func cleanupRoot() async { await self.puncherActor.stop() await self.arpResolver.stop() await self.sessionManager.clear() - await self.stunRequestWorker?.stop() - self.stunRequestWorker = nil - - await self.updatePolicyWorker?.stop() - self.updatePolicyWorker = nil await self.policyService.clear() await self.udpHoleService.stop() @@ -440,8 +454,6 @@ extension SDLContextActor { do { try await self.tunNetworkManager.apply(settings: .init(config: self.config), dnsServer: DNSHelper.dnsServer) SDLLogger.log("[SDLContext] setNetworkSettings successed") - // 开启权限的定时更新 - await self.whenRegistedSuper() await self.readySignal.succeed(()) } catch let err { SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)") @@ -450,34 +462,6 @@ extension SDLContextActor { } } - // 注册成功super的回调函数 - private func whenRegistedSuper() async { - await self.updatePolicyWorker?.stop() - let policyService = self.policyService - let superService = self.superService - - let updatePolicyWorker = PeriodicWorker( - configuration: .init( - interval: .seconds(10), - runImmediately: true, - mode: .fixedDelay, - errorPolicy: .keepRunning(delay: .seconds(5)) - ), - operation: { - SDLLogger.log("[SDLContext] updatePolicyTask execute") - await policyService.updatePolicy(superService: superService) - }, - onError: { err in - SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)") - } - ) - self.updatePolicyWorker = updatePolicyWorker - await updatePolicyWorker.start() - - // 启动stun任务 - await self.startStunRequestTask() - } - private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) async { let errorMessage = nakPacket.errorMessage guard let errorCode = SDLNAKErrorCode(rawValue: UInt8(nakPacket.errorCode)) else { @@ -641,35 +625,16 @@ extension SDLContextActor { // MARK: 和Stun相关的心跳机制 extension SDLContextActor { - - // MARK: -- StunRequestTask - private func startStunRequestTask() async { - await self.stunRequestWorker?.stop() + private func runStunRequestOnce() async throws { + let probeReply = try? await self.ipv6AssistClient?.probe(requestTimeout: .seconds(3)) - let stunRequestWorker = PeriodicWorker( - configuration: .init( - interval: .seconds(8), - runImmediately: true, - mode: .fixedDelay, - errorPolicy: .keepRunning(delay: .seconds(5)) - ), - operation: { [weak self] in - let probeReply = try? await self?.ipv6AssistClient?.probe(requestTimeout: .seconds(3)) + if let v6Info = probeReply?.v6Info, let v6Address = SDLUtil.ipv6DataToString(v6Info.v6) { + SDLLogger.log("[SDLContext] probe ipv6 address: \(v6Address)") + } else { + SDLLogger.log("[SDLContext] probe ipv6 address: empty") + } - if let v6Info = probeReply?.v6Info, let v6Address = SDLUtil.ipv6DataToString(v6Info.v6) { - SDLLogger.log("[SDLContext] probe ipv6 address: \(v6Address)") - } else { - SDLLogger.log("[SDLContext] probe ipv6 address: empty") - } - - await self?.sendStunRequest(v6Info: probeReply?.v6Info) - }, - onError: { err in - SDLLogger.log("[SDLContext] udp stunRequestTask stop with err: \(err)") - } - ) - self.stunRequestWorker = stunRequestWorker - await stunRequestWorker.start() + await self.sendStunRequest(v6Info: probeReply?.v6Info) } private func sendStunRequest(v6Info: SDLV6Info?) async { diff --git a/Tun/DNS/DNSService.swift b/Tun/DNS/DNSService.swift index 5647217..ed18068 100644 --- a/Tun/DNS/DNSService.swift +++ b/Tun/DNS/DNSService.swift @@ -13,10 +13,7 @@ actor DNSService { private let onEvent: EventHandler private var dnsClient: DNSCloudClient? - private var dnsMonitorTask: Task? - private var dnsLocalClient: DNSLocalClient? - private var dnsLocalMonitorTask: Task? init(serverIP: String, publicDnsServers: [String], onEvent: @escaping EventHandler) { self.serverIP = serverIP @@ -24,33 +21,37 @@ actor DNSService { self.onEvent = onEvent } - func start() { - self.startCloud() - self.startLocal() + func run() async throws { + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + group.addTask { + try await Self.runRestarting(name: "dnsServiceCloud") { + try await self.runCloud() + } + } + + group.addTask { + try await Self.runRestarting(name: "dnsServiceLocal") { + try await self.runLocal() + } + } + + try await group.waitForAll() + } } func stop() async { let dnsClient = self.dnsClient self.dnsClient = nil - let dnsMonitorTask = self.dnsMonitorTask - self.dnsMonitorTask = nil let dnsLocalClient = self.dnsLocalClient self.dnsLocalClient = nil - let dnsLocalMonitorTask = self.dnsLocalMonitorTask - self.dnsLocalMonitorTask = nil - dnsMonitorTask?.cancel() dnsClient?.stop() - if let dnsMonitorTask { - await dnsMonitorTask.value - } - - dnsLocalMonitorTask?.cancel() await dnsLocalClient?.stop() - if let dnsLocalMonitorTask { - await dnsLocalMonitorTask.value - } } func forward(ipPacketData: Data) { @@ -61,32 +62,6 @@ actor DNSService { await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: dnsPayload) } - private func startCloud() { - guard self.dnsMonitorTask == nil else { - return - } - - self.dnsMonitorTask = startMonitorTask(name: "dnsServiceCloudMonitor") { [weak self] in - guard let self else { - throw CancellationError() - } - try await self.runCloud() - } - } - - private func startLocal() { - guard self.dnsLocalMonitorTask == nil else { - return - } - - self.dnsLocalMonitorTask = startMonitorTask(name: "dnsServiceLocalMonitor") { [weak self] in - guard let self else { - throw CancellationError() - } - try await self.runLocal() - } - } - private func runCloud() async throws { let dnsClient = DNSCloudClient(serverIP: self.serverIP, port: 15353) self.dnsClient = dnsClient @@ -141,4 +116,25 @@ actor DNSService { throw error } } + + private static func runRestarting( + name: String, + retryDelay: Duration = .seconds(5), + operation: @escaping @Sendable () async throws -> Void + ) async throws { + while !Task.isCancelled { + do { + try Task.checkCancellation() + try await operation() + SDLLogger.log("[DNSService] worker \(name) ended, will restart", for: .debug) + } catch is CancellationError { + SDLLogger.log("[DNSService] worker \(name) cancelled", for: .debug) + throw CancellationError() + } catch { + SDLLogger.log("[DNSService] worker \(name) crashed: \(error.localizedDescription), will restart", for: .debug) + } + + try await Task.sleep(for: retryDelay) + } + } } diff --git a/Tun/Session/SDLPuncherActor.swift b/Tun/Session/SDLPuncherActor.swift index 7ac801c..b2ce6cf 100644 --- a/Tun/Session/SDLPuncherActor.swift +++ b/Tun/Session/SDLPuncherActor.swift @@ -49,18 +49,12 @@ actor SDLPuncherActor { // dstMac private var requestEntries: [Data: RequestEntry] = [:] - private var cleanupTask: Task? - func start() { - guard self.cleanupTask == nil else { - return - } - - self.cleanupTask = Task { [weak self] in - while !Task.isCancelled { - try? await Task.sleep(for: .seconds(1)) - await self?.cleanupExpiredEntries() - } + func runCleanup() async throws { + while !Task.isCancelled { + try await Task.sleep(for: .seconds(1)) + try Task.checkCancellation() + self.cleanupExpiredEntries() } } @@ -122,8 +116,6 @@ actor SDLPuncherActor { } func stop() { - self.cleanupTask?.cancel() - self.cleanupTask = nil self.requestEntries.removeAll() } @@ -133,7 +125,4 @@ actor SDLPuncherActor { } } - deinit { - self.cleanupTask?.cancel() - } }