From 8183e3b3bf9886dae42cd4353c00911c75d13996 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Wed, 27 May 2026 15:17:05 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E6=B5=81=E7=A8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Tun/App/SDLRuntimeEnvironment.swift | 2 +- Tun/Context/SDLContextActor.swift | 167 ++++++++++--- Tun/Inbound/PacketInboundActor.swift | 8 +- Tun/Outbound/PacketOutboundActor.swift | 18 +- Tun/Policy/PolicyService.swift | 4 +- Tun/Super/SDLSuperService.swift | 185 +++++++++----- Tun/UDPHole/SDLUDPHoleService.swift | 331 +++++++++++++++---------- 7 files changed, 463 insertions(+), 252 deletions(-) diff --git a/Tun/App/SDLRuntimeEnvironment.swift b/Tun/App/SDLRuntimeEnvironment.swift index 74e476d..482e12e 100644 --- a/Tun/App/SDLRuntimeEnvironment.swift +++ b/Tun/App/SDLRuntimeEnvironment.swift @@ -83,7 +83,7 @@ final class SDLRuntimeEnvironment { ) self.contextActor = contextActor - await contextActor.start() + try await contextActor.start() self.state = .running case .running: SDLLogger.log("[SDLRuntimeEnvironment] is running, ignore start command") diff --git a/Tun/Context/SDLContextActor.swift b/Tun/Context/SDLContextActor.swift index ccca43a..b3b8268 100644 --- a/Tun/Context/SDLContextActor.swift +++ b/Tun/Context/SDLContextActor.swift @@ -64,8 +64,8 @@ actor SDLContextActor { nonisolated let rsaCipher: RSACipher private var dnsService: DNSService? - private let superServiceProxy: SDLSuperServiceProxy - private let udpHoleServiceProxy: SDLUDPHoleServiceProxy + private let superService: SDLSuperService + private let udpHoleService: SDLUDPHoleService private let packetOutboundActor: PacketOutboundActor private let packetInboundActor: PacketInboundActor private let tunNetworkManager: SDLTunNetworkManager @@ -96,6 +96,8 @@ actor SDLContextActor { // stunRequest任务 private var stunRequestWorker: PeriodicWorker? + private var rootTask: Task? + private let readySignal = AsyncOneShot() public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) { let puncherActor = SDLPuncherActor() @@ -104,8 +106,8 @@ actor SDLContextActor { let arpResolver = ArpResolver() let flowTracer = SDLFlowTracer() let policyService = PolicyService(identityId: config.identityId, acl: config.acl) - let superServiceProxy = SDLSuperServiceProxy() - let udpHoleServiceProxy = SDLUDPHoleServiceProxy() + let superService = SDLSuperService(serverEndpoint: config.serverEndpoint) + let udpHoleService = SDLUDPHoleService(proberActor: proberActor) let tunNetworkManager = SDLTunNetworkManager(provider: provider) let packetOutboundActor = PacketOutboundActor( provider: provider, @@ -115,8 +117,8 @@ actor SDLContextActor { arpResolver: arpResolver, puncherActor: puncherActor, policyService: policyService, - superServiceProxy: superServiceProxy, - udpHoleServiceProxy: udpHoleServiceProxy, + superService: superService, + udpHoleService: udpHoleService, flowTracer: flowTracer ) let packetInboundActor = PacketInboundActor( @@ -126,7 +128,7 @@ actor SDLContextActor { policyService: policyService, packetOutboundActor: packetOutboundActor, arpResolver: arpResolver, - superServiceProxy: superServiceProxy, + superService: superService, flowTracer: flowTracer ) @@ -143,20 +145,66 @@ actor SDLContextActor { // 权限控制 self.policyService = policyService - self.superServiceProxy = superServiceProxy - self.udpHoleServiceProxy = udpHoleServiceProxy + self.superService = superService + self.udpHoleService = udpHoleService self.packetOutboundActor = packetOutboundActor self.packetInboundActor = packetInboundActor self.tunNetworkManager = tunNetworkManager } - public func start() async { + public func start() async throws { + guard self.rootTask == nil else { + try await self.readySignal.wait(timeout: .seconds(30)) + return + } + + let rootTask = Task { + try await self.runRoot() + } + self.rootTask = rootTask + + do { + try await self.readySignal.wait(timeout: .seconds(30)) + } catch { + rootTask.cancel() + _ = try? await rootTask.value + self.rootTask = nil + throw error + } + } + + // 处理context的停止问题 + public func stop() async { + let rootTask = self.rootTask + self.rootTask = nil + + rootTask?.cancel() + _ = try? await rootTask?.value + + await self.cleanupRoot() + } + + private func runRoot() async throws { + do { + try await self.runRootBody() + await self.cleanupRoot() + } catch is CancellationError { + await self.cleanupRoot() + throw CancellationError() + } catch { + await self.readySignal.fail(error) + await self.cleanupRoot() + throw error + } + } + + 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) } @@ -164,29 +212,66 @@ actor SDLContextActor { await self.packetOutboundActor.updateDNSService(dnsService) await dnsService.start() - let udpHoleEventHandler = await self.udpHoleServiceProxy.makeEventHandler { [weak self] event in - await self?.handleUDPHoleControlEvent(event) + await self.superService.updateMessageHandler { [weak self] message in + await self?.handleSuperMessage(message: message) } + let packetInboundActor = self.packetInboundActor - let udpHoleService = SDLUDPHoleService( - proberActor: self.proberActor, - onEvent: udpHoleEventHandler, + await self.udpHoleService.updateHandlers( + onEvent: { [weak self] event in + await self?.handleUDPHoleControlEvent(event) + }, onData: { data in await packetInboundActor.handleData(data) } ) - await self.udpHoleServiceProxy.replace(udpHoleService) - await udpHoleService.start(includeV6: false) - let superService = SDLSuperService(serverEndpoint: self.config.serverEndpoint) { [weak self] message in - await self?.handleSuperMessage(message: message) + let superService = self.superService + let udpHoleService = self.udpHoleService + + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + group.addTask { + try await Self.runRestarting(name: "superService") { + try await superService.run() + } + } + + group.addTask { + try await Self.runRestarting(name: "udpHoleService") { + try await udpHoleService.run(includeV6: false) + } + } + + try await group.waitForAll() } - await self.superServiceProxy.replace(superService) - await superService.start() } - // 处理context的停止问题 - public func stop() async { + 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("[SDLContext] worker \(name) ended, will restart", for: .debug) + } catch is CancellationError { + SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug) + throw CancellationError() + } catch { + SDLLogger.log("[SDLContext] worker \(name) crashed: \(error.localizedDescription), will restart", for: .debug) + } + + try await Task.sleep(for: retryDelay) + } + } + + private func cleanupRoot() async { await self.puncherActor.stop() await self.arpResolver.stop() await self.sessionManager.clear() @@ -197,24 +282,23 @@ actor SDLContextActor { await self.updatePolicyWorker?.stop() self.updatePolicyWorker = nil await self.policyService.clear() - - await self.packetOutboundActor.stop() - await self.udpHoleServiceProxy.stop() + await self.packetOutboundActor.stop() + await self.udpHoleService.stop() let dnsService = self.dnsService self.dnsService = nil await self.packetOutboundActor.updateDNSService(nil) await dnsService?.stop() - await self.superServiceProxy.stop() - + await self.superService.stop() + self.sessionToken = nil self.dataCipher = nil self.natType = .blocked await self.packetOutboundActor.updateRuntime(config: self.config, dataCipher: nil) await self.packetInboundActor.updateRuntime(config: self.config, dataCipher: nil) - + await self.ipv6AssistClient?.stop() self.ipv6AssistClient = nil } @@ -263,7 +347,7 @@ extension SDLContextActor { } private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { - await self.udpHoleServiceProxy.send(type: type, data: data, remoteAddress: remoteAddress) + await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) } } @@ -292,12 +376,12 @@ extension SDLContextActor { case .registerSuperAck(let registerSuperAck): await self.handleRegisterSuperAck(registerSuperAck: registerSuperAck) case .registerSuperNak(let registerSuperNak): - self.handleRegisterSuperNak(nakPacket: registerSuperNak) + await self.handleRegisterSuperNak(nakPacket: registerSuperNak) case .peerInfo(let peerInfo): SDLLogger.log("[SDLContext] peer message: \(peerInfo)") let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo) for packet in packets { - await self.udpHoleServiceProxy.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress) + await self.udpHoleService.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress) } case .event(let event): await self.handleEvent(event: event) @@ -321,6 +405,7 @@ extension SDLContextActor { guard let key = try? self.rsaCipher.decode(data: Data(registerSuperAck.key)) else { SDLLogger.log("[SDLContext] registerSuperAck invalid key") let error = SDLError.invalidKey + await self.readySignal.fail(error) self.provider.cancelTunnelWithError(error) return } @@ -337,6 +422,7 @@ extension SDLContextActor { default: SDLLogger.log("[SDLContext] registerSuperAck invalid algorithm \(algorithm)") let error = SDLError.unsupportedAlgorithm(algorithm: algorithm) + await self.readySignal.fail(error) self.provider.cancelTunnelWithError(error) return } @@ -351,8 +437,10 @@ extension SDLContextActor { await self.packetOutboundActor.startPacketReader() // 开启权限的定时更新 await self.whenRegistedSuper() + await self.readySignal.succeed(()) } catch let err { SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)") + await self.readySignal.fail(err) self.provider.cancelTunnelWithError(err) } } @@ -361,7 +449,7 @@ extension SDLContextActor { private func whenRegistedSuper() async { await self.updatePolicyWorker?.stop() let policyService = self.policyService - let superServiceProxy = self.superServiceProxy + let superService = self.superService let updatePolicyWorker = PeriodicWorker( configuration: .init( @@ -372,7 +460,7 @@ extension SDLContextActor { ), operation: { SDLLogger.log("[SDLContext] updatePolicyTask execute") - await policyService.updatePolicy(superServiceProxy: superServiceProxy) + await policyService.updatePolicy(superService: superService) }, onError: { err in SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)") @@ -385,7 +473,7 @@ extension SDLContextActor { await self.startStunRequestTask() } - private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) { + private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) async { let errorMessage = nakPacket.errorMessage guard let errorCode = SDLNAKErrorCode(rawValue: UInt8(nakPacket.errorCode)) else { return @@ -396,6 +484,7 @@ extension SDLContextActor { self.publishTunnelEvent(code: Int(errorCode.rawValue), message: errorMessage) // 报告错误并退出 let error = NSError(domain: "com.jihe.punchnet.tun", code: -1) + await self.readySignal.fail(error) self.provider.cancelTunnelWithError(error) case .noIpAddress, .networkFault, .internalFault: @@ -445,7 +534,7 @@ extension SDLContextActor { if let registerSuperData = try? registerSuper.serializedData() { SDLLogger.log("[SDLContext] will send register super") - await self.superServiceProxy.send(type: .registerSuper, data: registerSuperData) + await self.superService.send(type: .registerSuper, data: registerSuperData) } } @@ -454,7 +543,7 @@ extension SDLContextActor { return } - await self.superServiceProxy.send(type: .exposedServiceRequest, data: requestData) + await self.superService.send(type: .exposedServiceRequest, data: requestData) } private func applyExposedServiceResponse(_ response: SDLExposedServiceResponse) async { diff --git a/Tun/Inbound/PacketInboundActor.swift b/Tun/Inbound/PacketInboundActor.swift index f33d47c..38bbee9 100644 --- a/Tun/Inbound/PacketInboundActor.swift +++ b/Tun/Inbound/PacketInboundActor.swift @@ -42,7 +42,7 @@ actor PacketInboundActor { private let policyService: PolicyService private let packetOutboundActor: PacketOutboundActor private let arpResolver: ArpResolver - private let superServiceProxy: SDLSuperServiceProxy + private let superService: SDLSuperService private let flowTracer: SDLFlowTracer private var networkAddress: SDLConfiguration.NetworkAddress @@ -55,7 +55,7 @@ actor PacketInboundActor { policyService: PolicyService, packetOutboundActor: PacketOutboundActor, arpResolver: ArpResolver, - superServiceProxy: SDLSuperServiceProxy, + superService: SDLSuperService, flowTracer: SDLFlowTracer) { self.provider = provider self.networkAddress = config.networkAddress @@ -64,7 +64,7 @@ actor PacketInboundActor { self.policyService = policyService self.packetOutboundActor = packetOutboundActor self.arpResolver = arpResolver - self.superServiceProxy = superServiceProxy + self.superService = superService self.flowTracer = flowTracer } @@ -96,7 +96,7 @@ actor PacketInboundActor { case .requestPolicy(let context): SDLLogger.log("[PacketInboundActor] policy miss, \(context.logDescription)", for: .debug) if let queryData = await self.policyService.makePolicyRequest(srcIdentityID: context.srcIdentityID) { - await self.superServiceProxy.send(type: .policyRequest, data: queryData) + await self.superService.send(type: .policyRequest, data: queryData) } case .dropByPolicy(let context): SDLLogger.log("[PacketInboundActor] policy denied, \(context.logDescription)", for: .trace) diff --git a/Tun/Outbound/PacketOutboundActor.swift b/Tun/Outbound/PacketOutboundActor.swift index 54b6c1b..4ca2191 100644 --- a/Tun/Outbound/PacketOutboundActor.swift +++ b/Tun/Outbound/PacketOutboundActor.swift @@ -23,8 +23,8 @@ actor PacketOutboundActor { private let arpResolver: ArpResolver private let puncherActor: SDLPuncherActor private let policyService: PolicyService - private let superServiceProxy: SDLSuperServiceProxy - private let udpHoleServiceProxy: SDLUDPHoleServiceProxy + private let superService: SDLSuperService + private let udpHoleService: SDLUDPHoleService private let flowTracer: SDLFlowTracer private var packetReaderTask: Task? private var packetReaderGeneration: UInt64 = 0 @@ -43,8 +43,8 @@ actor PacketOutboundActor { arpResolver: ArpResolver, puncherActor: SDLPuncherActor, policyService: PolicyService, - superServiceProxy: SDLSuperServiceProxy, - udpHoleServiceProxy: SDLUDPHoleServiceProxy, + superService: SDLSuperService, + udpHoleService: SDLUDPHoleService, flowTracer: SDLFlowTracer) { self.provider = provider self.networkAddress = config.networkAddress @@ -56,8 +56,8 @@ actor PacketOutboundActor { self.arpResolver = arpResolver self.puncherActor = puncherActor self.policyService = policyService - self.superServiceProxy = superServiceProxy - self.udpHoleServiceProxy = udpHoleServiceProxy + self.superService = superService + self.udpHoleService = udpHoleService self.flowTracer = flowTracer } @@ -159,7 +159,7 @@ actor PacketOutboundActor { } else { SDLLogger.log("[PacketOutboundActor] dstIp: \(SDLUtil.int32ToIp(ip)) arp query not found, broadcast", for: .trace) if let arpRequest = try? await self.arpResolver.makeArpRequest(targetIp: ip) { - await self.superServiceProxy.send(type: .arpRequest, data: arpRequest) + await self.superService.send(type: .arpRequest, data: arpRequest) } } } @@ -183,7 +183,7 @@ actor PacketOutboundActor { self.flowTracer.inc(num: payload.count, type: .forward) if let queryData = await self.puncherActor.makeQueryInfoRequest(request: request) { - await self.superServiceProxy.send(type: .queryInfo, data: queryData) + await self.superService.send(type: .queryInfo, data: queryData) } } @@ -265,6 +265,6 @@ actor PacketOutboundActor { } private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { - await self.udpHoleServiceProxy.send(type: type, data: data, remoteAddress: remoteAddress) + await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) } } diff --git a/Tun/Policy/PolicyService.swift b/Tun/Policy/PolicyService.swift index 3721312..1b10c4b 100644 --- a/Tun/Policy/PolicyService.swift +++ b/Tun/Policy/PolicyService.swift @@ -59,10 +59,10 @@ actor PolicyService { return await self.policyRuleStore.makePolicyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.identityId) } - func updatePolicy(superServiceProxy: SDLSuperServiceProxy) async { + func updatePolicy(superService: SDLSuperService) async { let requests = await self.policyRuleStore.makeBatchPolicyRequests(dstIdentityID: self.identityId) for request in requests { - await superServiceProxy.send(type: .policyRequest, data: request) + await superService.send(type: .policyRequest, data: request) } } diff --git a/Tun/Super/SDLSuperService.swift b/Tun/Super/SDLSuperService.swift index 52cd4bb..78406ba 100644 --- a/Tun/Super/SDLSuperService.swift +++ b/Tun/Super/SDLSuperService.swift @@ -5,109 +5,160 @@ actor SDLSuperService { private let serverEndpoint: SDLConfiguration.ResolvedServerEndpoint private let port: UInt16 - private let onMessage: MessageHandler - private var superClient: SDLSuperClient? - private var monitorTask: Task? + private var onMessage: MessageHandler = { _ in } + private var currentSession: SDLSuperSession? + private var generation: UInt64 = 0 - init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16 = 1443, onMessage: @escaping MessageHandler) { + init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16 = 1443) { self.serverEndpoint = serverEndpoint self.port = port + } + + func updateMessageHandler(_ onMessage: @escaping MessageHandler) { self.onMessage = onMessage } - func start() { - guard self.monitorTask == nil else { - return - } - - self.monitorTask = startMonitorTask(name: "superServiceMonitor") { [weak self] in - guard let self else { - throw CancellationError() + func run() async throws { + let generation = self.nextGeneration() + let session = SDLSuperSession( + serverEndpoint: self.serverEndpoint, + port: self.port, + onMessage: { [weak self] message in + await self?.handleMessage(message, generation: generation) } - try await self.runOnce() - } - } + ) - func stop() async { - let monitorTask = self.monitorTask - self.monitorTask = nil - - let superClient = self.superClient - self.superClient = nil - - monitorTask?.cancel() - await superClient?.stop() - - if let monitorTask { - await monitorTask.value - } - } - - func send(type: SDLPacketType, data: Data) async { - await self.superClient?.send(type: type, data: data) - } - - private func runOnce() async throws { - let superClient = SDLSuperClient(serverEndpoint: self.serverEndpoint, port: self.port) - self.superClient = superClient - await superClient.start() + self.currentSession = session do { - try await withTaskCancellationHandler { - try await self.run(superClient) - } onCancel: { - Task { - await superClient.stop() - } - } - - await self.cleanup(superClient) + try await session.run() + self.clearCurrent(session, generation: generation) + } catch is CancellationError { + self.clearCurrent(session, generation: generation) + await session.stop() + throw CancellationError() } catch { - await self.cleanup(superClient) + self.clearCurrent(session, generation: generation) + await session.stop() throw error } } - private func run(_ superClient: SDLSuperClient) async throws { + func stop() async { + self.generation &+= 1 + + let session = self.currentSession + self.currentSession = nil + + await session?.stop() + } + + func send(type: SDLPacketType, data: Data) async { + await self.currentSession?.send(type: type, data: data) + } + + private func nextGeneration() -> UInt64 { + self.generation &+= 1 + return self.generation + } + + private func clearCurrent(_ session: SDLSuperSession, generation: UInt64) { + guard self.generation == generation else { + return + } + + if self.currentSession === session { + self.currentSession = nil + } + } + + private func handleMessage(_ message: SDLQUICInboundMessage, generation: UInt64) async { + guard self.generation == generation else { + return + } + + await self.onMessage(message) + } +} + +final class SDLSuperSession: @unchecked Sendable { + typealias MessageHandler = @Sendable (SDLQUICInboundMessage) async -> Void + + private let serverEndpoint: SDLConfiguration.ResolvedServerEndpoint + private let port: UInt16 + private let onMessage: MessageHandler + private let client: SDLSuperClient + + init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16, onMessage: @escaping MessageHandler) { + self.serverEndpoint = serverEndpoint + self.port = port + self.onMessage = onMessage + self.client = SDLSuperClient(serverEndpoint: serverEndpoint, port: port) + } + + func run() async throws { + await self.client.start() + + do { + try await withTaskCancellationHandler { + try await self.runLoops() + } onCancel: { + Task { + await self.client.stop() + } + } + + await self.stop() + } catch { + await self.stop() + throw error + } + } + + func stop() async { + await self.client.stop() + } + + func send(type: SDLPacketType, data: Data) async { + await self.client.send(type: type, data: data) + } + + private func runLoops() async throws { try await Task.sleep(for: .seconds(0.5)) try Task.checkCancellation() - SDLLogger.log("[SDLSuperService] start super client: \(self.serverEndpoint.ip)") + SDLLogger.log("[SDLSuperSession] start super client: \(self.serverEndpoint.ip)") try await withThrowingTaskGroup(of: Void.self) { group in defer { group.cancelAll() } - let onMessage = self.onMessage group.addTask { - for try await message in superClient.messageStream { - try Task.checkCancellation() - await onMessage(message) - } + try await self.readLoop() } group.addTask { - while true { - try await Task.sleep(for: .seconds(5)) - try Task.checkCancellation() - await superClient.send(type: .ping, data: Data()) - } + try await self.pingLoop() } _ = try await group.next() } } - private func cleanup(_ superClient: SDLSuperClient) async { - await superClient.stop() - - if self.superClient === superClient { - self.superClient = nil + private func readLoop() async throws { + for try await message in self.client.messageStream { + try Task.checkCancellation() + await self.onMessage(message) } + } - SDLLogger.log("[SDLSuperService] cleanup") + private func pingLoop() async throws { + while true { + try await Task.sleep(for: .seconds(5)) + try Task.checkCancellation() + await self.client.send(type: .ping, data: Data()) + } } } - diff --git a/Tun/UDPHole/SDLUDPHoleService.swift b/Tun/UDPHole/SDLUDPHoleService.swift index a1db8fd..20b156c 100644 --- a/Tun/UDPHole/SDLUDPHoleService.swift +++ b/Tun/UDPHole/SDLUDPHoleService.swift @@ -27,27 +27,131 @@ actor SDLUDPHoleService { typealias DataHandler = @Sendable (SDLData) async -> Void private let proberActor: SDLNATProberActor - private let onEvent: EventHandler - private let onData: DataHandler - private var udpHole: SDLUDPHole? - private var udpHoleMonitorTask: Task? - private var natProbeTask: Task? - private var localAddress: SocketAddress? + private var onEvent: EventHandler = { _ in } + private var onData: DataHandler = { _ in } + private var currentSession: SDLUDPHoleSession? + private var generation: UInt64 = 0 - private var udpHoleV6: SDLUDPHoleV6? - private var udpHoleV6MonitorTask: Task? - - init(proberActor: SDLNATProberActor, onEvent: @escaping EventHandler, onData: @escaping DataHandler) { + init(proberActor: SDLNATProberActor) { self.proberActor = proberActor + } + + func updateHandlers(onEvent: @escaping EventHandler, onData: @escaping DataHandler) { self.onEvent = onEvent self.onData = onData } - func start(includeV6: Bool = false) { - self.startV4() - if includeV6 { - self.startV6() + func run(includeV6: Bool = false) async throws { + let generation = self.nextGeneration() + let session = SDLUDPHoleSession( + proberActor: self.proberActor, + includeV6: includeV6, + onEvent: { [weak self] event in + await self?.handleEvent(event, generation: generation) + }, + onData: self.onData + ) + + self.currentSession = session + + do { + try await session.run() + self.clearCurrent(session, generation: generation) + } catch is CancellationError { + self.clearCurrent(session, generation: generation) + await session.stop() + throw CancellationError() + } catch { + self.clearCurrent(session, generation: generation) + await session.stop() + throw error + } + } + + func stop() async { + self.generation &+= 1 + + let session = self.currentSession + self.currentSession = nil + + await session?.stop() + await self.proberActor.cancelAll() + } + + func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { + await self.currentSession?.send(type: type, data: data, remoteAddress: remoteAddress) + } + + private func nextGeneration() -> UInt64 { + self.generation &+= 1 + return self.generation + } + + private func clearCurrent(_ session: SDLUDPHoleSession, generation: UInt64) { + guard self.generation == generation else { + return + } + + if self.currentSession === session { + self.currentSession = nil + } + } + + private func handleEvent(_ event: Event, generation: UInt64) async { + guard self.generation == generation else { + return + } + + await self.onEvent(event) + } +} + +actor SDLUDPHoleSession { + private let proberActor: SDLNATProberActor + private let includeV6: Bool + private let onEvent: SDLUDPHoleService.EventHandler + private let onData: SDLUDPHoleService.DataHandler + + private var udpHole: SDLUDPHole? + private var udpHoleV6: SDLUDPHoleV6? + private var localAddress: SocketAddress? + + init( + proberActor: SDLNATProberActor, + includeV6: Bool, + onEvent: @escaping SDLUDPHoleService.EventHandler, + onData: @escaping SDLUDPHoleService.DataHandler + ) { + self.proberActor = proberActor + self.includeV6 = includeV6 + self.onEvent = onEvent + self.onData = onData + } + + func run() async throws { + do { + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + group.addTask { + try await self.runV4() + } + + if self.includeV6 { + group.addTask { + try await self.runV6() + } + } + + try await group.waitForAll() + } + await self.stop() + } catch { + await self.stop() + throw error } } @@ -56,64 +160,32 @@ actor SDLUDPHoleService { self.udpHole = nil self.localAddress = nil - let udpHoleMonitorTask = self.udpHoleMonitorTask - self.udpHoleMonitorTask = nil - - let natProbeTask = self.natProbeTask - self.natProbeTask = nil - - udpHoleMonitorTask?.cancel() - natProbeTask?.cancel() - await self.proberActor.cancelAll() - await udpHole?.stop() - - if let natProbeTask { - await natProbeTask.value - } - if let udpHoleMonitorTask { - await udpHoleMonitorTask.value - } - let udpHoleV6 = self.udpHoleV6 self.udpHoleV6 = nil - let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask - self.udpHoleV6MonitorTask = nil - udpHoleV6MonitorTask?.cancel() + + await self.proberActor.cancelAll() + await udpHole?.stop() udpHoleV6?.stop() - if let udpHoleV6MonitorTask { - await udpHoleV6MonitorTask.value - } } func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { switch remoteAddress { case .v4: guard let udpHole else { - SDLLogger.log("[SDLUDPHoleService] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) + SDLLogger.log("[SDLUDPHoleSession] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) return } + await udpHole.send(type: type, data: data, remoteAddress: remoteAddress) case .v6: guard let udpHoleV6 else { - SDLLogger.log("[SDLUDPHoleService] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) + SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) return } + udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress) default: - SDLLogger.log("[SDLUDPHoleService] unsupported socket family: \(remoteAddress)", for: .debug) - } - } - - private func startV4() { - guard self.udpHoleMonitorTask == nil else { - return - } - - self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleServiceV4Monitor") { [weak self] in - guard let self else { - throw CancellationError() - } - try await self.runV4() + SDLLogger.log("[SDLUDPHoleSession] unsupported socket family: \(remoteAddress)", for: .debug) } } @@ -122,23 +194,26 @@ actor SDLUDPHoleService { let localAddress = try await udpHole.start() self.udpHole = udpHole self.localAddress = localAddress - SDLLogger.log("[SDLUDPHoleService] udpHole started, on address: \(localAddress)") + SDLLogger.log("[SDLUDPHoleSession] udpHole started, on address: \(localAddress)") await self.onEvent(.ready(localAddress)) - self.startNatProbe(using: udpHole) - - defer { - if self.udpHole === udpHole { - self.udpHole = nil - self.localAddress = nil - } - } do { try await withTaskCancellationHandler { - for try await datagram in await udpHole.messageStream() { - try Task.checkCancellation() - try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message) + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + group.addTask { + try await self.readV4Loop(udpHole: udpHole) + } + + group.addTask { + await self.probeNatType(udpHole: udpHole) + } + + try await group.waitForAll() } } onCancel: { Task { @@ -147,26 +222,34 @@ actor SDLUDPHoleService { } } catch { await udpHole.stop() + if self.udpHole === udpHole { + self.udpHole = nil + self.localAddress = nil + } throw error } } - private func startNatProbe(using udpHole: SDLUDPHole) { - self.natProbeTask?.cancel() - let proberActor = self.proberActor - let onEvent = self.onEvent - self.natProbeTask = Task { - if Task.isCancelled { - return - } - let natType = await proberActor.probeNatType(using: udpHole) - if Task.isCancelled { - return - } - await onEvent(.natType(natType)) + private func readV4Loop(udpHole: SDLUDPHole) async throws { + for try await datagram in await udpHole.messageStream() { + try Task.checkCancellation() + try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message) } } + private func probeNatType(udpHole: SDLUDPHole) async { + if Task.isCancelled { + return + } + + let natType = await self.proberActor.probeNatType(using: udpHole) + if Task.isCancelled { + return + } + + await self.onEvent(.natType(natType)) + } + private func handleV4Message(remoteAddress: SocketAddress, message: SDLHoleMessage) async throws { switch message { case .control(let control): @@ -181,69 +264,57 @@ actor SDLUDPHoleService { } } - private func startV6() { - guard self.udpHoleV6MonitorTask == nil else { - return - } - - self.udpHoleV6MonitorTask = startMonitorTask(name: "udpHoleServiceV6Monitor") { [weak self] in - guard let self else { - throw CancellationError() - } - try await self.runV6() - } - } - private func runV6() async throws { let udpHoleV6 = try SDLUDPHoleV6() let localAddress = try udpHoleV6.start() self.udpHoleV6 = udpHoleV6 if let localAddress { - SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, on address: \(localAddress)") + SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, on address: \(localAddress)") } else { - SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, no local address") + SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, no local address") } - defer { + do { + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + let onEvent = self.onEvent + let onData = self.onData + group.addTask { + for await (remoteAddress, message) in udpHoleV6.messageStream { + try Task.checkCancellation() + switch message { + case .control(let control): + await onEvent(.packet(remoteAddress, control, source: .v6)) + case .data(let data): + await onData(data) + } + } + } + + group.addTask { + for await event in udpHoleV6.eventStream { + try Task.checkCancellation() + switch event { + case .ready: + SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 ready") + case .closed, .errorCaught: + throw SDLContextError.udpHoleClosed + } + } + } + + _ = try await group.next() + } + } catch { + udpHoleV6.stop() if self.udpHoleV6 === udpHoleV6 { - udpHoleV6.stop() self.udpHoleV6 = nil } - } - - try await withThrowingTaskGroup(of: Void.self) { group in - defer { - group.cancelAll() - } - - let onEvent = self.onEvent - let onData = self.onData - group.addTask { - for await (remoteAddress, message) in udpHoleV6.messageStream { - try Task.checkCancellation() - switch message { - case .control(let control): - await onEvent(.packet(remoteAddress, control, source: .v6)) - case .data(let data): - await onData(data) - } - } - } - - group.addTask { - for await event in udpHoleV6.eventStream { - try Task.checkCancellation() - switch event { - case .ready: - SDLLogger.log("[SDLUDPHoleService] udpHoleV6 ready") - case .closed, .errorCaught: - throw SDLContextError.udpHoleClosed - } - } - } - - _ = try await group.next() + throw error } } }