From 55701085e3201d80dc299a98628d54a95cd36cad Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Mon, 4 May 2026 20:15:08 +0800 Subject: [PATCH] fix udpHole --- Tun/Punchnet/Actors/SDLContextActor.swift | 128 ++++++++++++---------- Tun/Punchnet/UDPHole/SDLUDPHole.swift | 12 +- Tun/Punchnet/UDPHole/SDLUDPHoleV6.swift | 90 +++++---------- 3 files changed, 106 insertions(+), 124 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 903633b..a59fdb9 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -53,8 +53,6 @@ actor SDLContextActor { private var udpHoleLocalAddress: SocketAddress? private var udpHoleV6: SDLUDPHoleV6? - private var udpHoleV6Workers: [Task]? - private var udpHoleV6LocalAddress: SocketAddress? // dns的client对象 private var dnsClient: DNSCloudClient? @@ -128,7 +126,7 @@ actor SDLContextActor { if resetNotifier { self.prepareTunnelNotifier() } - + // 启动arp的定时清理任务 await self.puncherActor.start() await self.arpServer.start() @@ -147,11 +145,10 @@ actor SDLContextActor { try await self.startQUICClient() SDLLogger.log("[SDLContext] superClient closed!!!!") } - + // await self.supervisor.addWorker(name: "udpHoleV6") { -// let udpHoleV6 = try await self.startUDPHoleV6() // SDLLogger.log("[SDLContext] udp v6 running!!!!") -// try await udpHoleV6.waitClose() +// try await self.startUDPHoleV6() // SDLLogger.log("[SDLContext] udp v6 closed!!!!") // } } @@ -357,6 +354,7 @@ actor SDLContextActor { self.udpHoleLocalAddress = localAddress defer { + self.udpHole?.stop() self.udpHole = nil self.udpHoleLocalAddress = nil } @@ -371,8 +369,17 @@ actor SDLContextActor { try Task.checkCancellation() switch message.inboundMessage { - case .control(let controlMessage): - await self.handleHoleControlMessage(controlMessage, localAddress: localAddress, remoteAddress: remoteAddress, source: .v4) + case .control(let message): + switch message { + case .stunReply(_): + SDLLogger.log("[SDLContext] get a stunReply", for: .debug) + case .stunProbeReply(let probeReply): + await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply) + case .register(let register): + try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v4) + case .registerAck(let registerAck): + await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6) + } case .data(let data): try? await self.handleHoleData(data: data) } @@ -399,26 +406,64 @@ actor SDLContextActor { } } - private func startUDPHoleV6() async throws -> SDLUDPHoleV6 { - self.udpHoleV6Workers?.forEach {$0.cancel()} - self.udpHoleV6Workers = nil - + private func startUDPHoleV6() async throws { // 启动udp服务器 let udpHoleV6 = try SDLUDPHoleV6() let localAddress = try udpHoleV6.start() - SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)") - - // 处理消息流 - let messageStream = udpHoleV6.messageStream - let messageTask = Task.detached { - await self.consumeUDPHoleMessages(stream: messageStream, localAddress: localAddress, source: .v6) + self.udpHoleV6 = udpHoleV6 + + if let localAddress { + SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)") + } else { + SDLLogger.log("[SDLContext] udpHoleV6 started, no local address") } - self.udpHoleV6 = udpHoleV6 - self.udpHoleV6LocalAddress = localAddress - self.udpHoleV6Workers = [messageTask] + defer { + self.udpHoleV6?.stop() + self.udpHoleV6 = nil + } - return udpHoleV6 + try await withThrowingTaskGroup { group in + defer { + group.cancelAll() + } + + // 处理消息流 + group.addTask { + for await (remoteAddress, message) in udpHoleV6.messageStream { + try Task.checkCancellation() + + switch message.inboundMessage { + case .control(let message): + switch message { + case .register(let register): + try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v6) + case .registerAck(let registerAck): + await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6) + default: + () + } + case .data(let data): + try? await self.handleHoleData(data: data) + } + } + } + + group.addTask { + for await event in udpHoleV6.eventStream { + try Task.checkCancellation() + + switch event { + case .ready: + SDLLogger.log("[SDLContext] udpHoleV6 ready") + case .closed, .errorCaught: + throw SDLContextError.udpHoleClosed + } + } + } + + try await group.next() + } } // 处理context的停止问题 @@ -435,14 +480,12 @@ actor SDLContextActor { self.flowSessionManager.clear() + self.udpHole?.stop() self.udpHole = nil self.udpHoleLocalAddress = nil - self.udpHoleV6Workers?.forEach { $0.cancel() } - self.udpHoleV6Workers = nil self.udpHoleV6?.stop() self.udpHoleV6 = nil - self.udpHoleV6LocalAddress = nil self.quicClient?.stop() self.quicClient = nil @@ -682,7 +725,6 @@ actor SDLContextActor { self.udpHole = nil self.udpHoleLocalAddress = nil self.udpHoleV6 = nil - self.udpHoleV6LocalAddress = nil self.dnsClient = nil } } @@ -800,40 +842,6 @@ extension SDLContextActor { // 处理从Hole收到的数据 extension SDLContextActor { - private func consumeUDPHoleMessages(stream: AsyncStream<(SocketAddress, SDLHoleMessage)>, localAddress: SocketAddress, source: UDPHoleKind) async { - for await (remoteAddress, message) in stream { - if Task.isCancelled { - break - } - - switch message.inboundMessage { - case .control(let controlMessage): - await self.handleHoleControlMessage(controlMessage, localAddress: localAddress, remoteAddress: remoteAddress, source: source) - case .data(let data): - try? await self.handleHoleData(data: data) - } - } - } - - private func handleHoleControlMessage(_ message: SDLHoleControlMessage, localAddress: SocketAddress, remoteAddress: SocketAddress, source: UDPHoleKind) async { - switch message { - case .stunReply(_): - guard source == .v4 else { - return - } - SDLLogger.log("[SDLContext] get a stunReply", for: .debug) - case .stunProbeReply(let probeReply): - guard source == .v4 else { - return - } - await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply) - case .register(let register): - try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: source) - case .registerAck(let registerAck): - await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: source) - } - } - private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: UDPHoleKind) async throws { let networkAddr = config.networkAddress SDLLogger.log("[SDLContext] register packet: \(register), network_address: \(networkAddr)") diff --git a/Tun/Punchnet/UDPHole/SDLUDPHole.swift b/Tun/Punchnet/UDPHole/SDLUDPHole.swift index 57a72d5..d258202 100644 --- a/Tun/Punchnet/UDPHole/SDLUDPHole.swift +++ b/Tun/Punchnet/UDPHole/SDLUDPHole.swift @@ -24,6 +24,8 @@ final class SDLUDPHole: ChannelInboundHandler { case errorCaught } + private var isStopped: Bool = false + private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) private var channel: Channel? @@ -113,8 +115,14 @@ final class SDLUDPHole: ChannelInboundHandler { } } - deinit { - SDLLogger.log("[SDLUDPHole] deinit", for: .debug) + func stop() { + guard !self.isStopped else { + return + } + + self.isStopped = true + + SDLLogger.log("[SDLUDPHole] stop", for: .debug) self.messageContinuation.finish() self.eventContinuation.finish() self.channel = nil diff --git a/Tun/Punchnet/UDPHole/SDLUDPHoleV6.swift b/Tun/Punchnet/UDPHole/SDLUDPHoleV6.swift index 907225d..d798e09 100644 --- a/Tun/Punchnet/UDPHole/SDLUDPHoleV6.swift +++ b/Tun/Punchnet/UDPHole/SDLUDPHoleV6.swift @@ -14,30 +14,37 @@ import SwiftProtobuf final class SDLUDPHoleV6: ChannelInboundHandler { typealias InboundIn = AddressedEnvelope - private enum State: Equatable { - case idle + // 事件 + enum HoleEvent { case ready - case stopping - case stopped + case closed + case errorCaught } + private var isStopped: Bool = false + private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) private var channel: Channel? - private var closeFuture: EventLoopFuture? - private var state: State = .idle - private var didFinishMessageStream: Bool = false public let messageStream: AsyncStream<(SocketAddress, SDLHoleMessage)> private let messageContinuation: AsyncStream<(SocketAddress, SDLHoleMessage)>.Continuation + // 事件相关逻辑 + public let eventStream: AsyncStream + private let eventContinuation: AsyncStream.Continuation + // 启动函数 init() throws { let (stream, continuation) = AsyncStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048)) self.messageStream = stream self.messageContinuation = continuation + + let eventPair = AsyncStream.makeStream(of: HoleEvent.self) + self.eventStream = eventPair.stream + self.eventContinuation = eventPair.continuation } - func start() throws -> SocketAddress { + func start() throws -> SocketAddress? { let bootstrap = DatagramBootstrap(group: group) .channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) .channelInitializer { channel in @@ -47,51 +54,13 @@ final class SDLUDPHoleV6: ChannelInboundHandler { // 绑定到IPv6通配地址,只处理IPv6流量 let channel = try bootstrap.bind(host: "::", port: 0).wait() self.channel = channel - self.closeFuture = channel.closeFuture - self.state = .ready - precondition(channel.localAddress != nil, "UDP v6 channel has no localAddress after bind") - return channel.localAddress! - } - - func waitClose() async throws { - switch self.state { - case .idle: - SDLLogger.log("[SDLUDPHoleV6] waitClose11", for: .debug) - return - case .ready, .stopping, .stopped: - guard let closeFuture = self.closeFuture else { - SDLLogger.log("[SDLUDPHoleV6] waitClose22", for: .debug) - return - } - try await closeFuture.get() - SDLLogger.log("[SDLUDPHoleV6] waitClose33", for: .debug) - } - } - - func stop() { - switch self.state { - case .stopping, .stopped: - return - case .idle: - self.state = .stopped - self.finishMessageStream() - return - case .ready: - self.state = .stopping - } - - self.finishMessageStream() - self.channel?.close(promise: nil) + return channel.localAddress } // --MARK: ChannelInboundHandler delegate func channelRead(context: ChannelHandlerContext, data: NIOAny) { - guard case .ready = self.state else { - return - } - let envelope = unwrapInboundIn(data) var buffer = envelope.data @@ -113,23 +82,19 @@ final class SDLUDPHoleV6: ChannelInboundHandler { } func channelInactive(context: ChannelHandlerContext) { - self.finishMessageStream() - self.channel = nil - self.state = .stopped + SDLLogger.log("[SDLUDPHoleV6] channelInactive", for: .debug) + self.eventContinuation.yield(.closed) } func errorCaught(context: ChannelHandlerContext, error: any Error) { SDLLogger.log("[SDLUDPHoleV6] channel error: \(error)", for: .debug) - self.finishMessageStream() - if self.state != .stopped { - self.state = .stopping - } context.close(promise: nil) + self.eventContinuation.yield(.errorCaught) } // MARK: 处理写入逻辑 func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { - guard case .ready = self.state, let channel = self.channel else { + guard let channel = self.channel else { return } @@ -143,17 +108,18 @@ final class SDLUDPHoleV6: ChannelInboundHandler { } } - private func finishMessageStream() { - guard !self.didFinishMessageStream else { + func stop() { + guard !self.isStopped else { return } - self.didFinishMessageStream = true + self.isStopped = true + SDLLogger.log("[SDLUDPHoleV6] stop", for: .debug) + self.messageContinuation.finish() - } - - deinit { - self.stop() + self.eventContinuation.finish() + self.channel = nil + try? self.group.syncShutdownGracefully() }