From 783c1329e13677103cb454daf87e21aaca01daa0 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Mon, 4 May 2026 19:51:07 +0800 Subject: [PATCH] fix udpHole --- Tun/Punchnet/Actors/SDLContextActor.swift | 74 ++++++++++++----- Tun/Punchnet/UDPHole/SDLUDPHole.swift | 98 +++++++---------------- 2 files changed, 82 insertions(+), 90 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index e38eafa..903633b 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -13,6 +13,11 @@ import NIOCore /* 1. 处理rsa的加解密逻辑 */ + +enum SDLContextError: Error { + case udpHoleClosed +} + actor SDLContextActor { private enum UDPHoleKind: Equatable { @@ -45,7 +50,6 @@ actor SDLContextActor { // 依赖的变量 private var udpHole: SDLUDPHole? - private var udpHoleWorkers: [Task]? private var udpHoleLocalAddress: SocketAddress? private var udpHoleV6: SDLUDPHoleV6? @@ -133,14 +137,15 @@ actor SDLContextActor { // 先启动udp await self.supervisor.addWorker(name: "udpHole") { - let udpHole = try await self.startUDPHole() SDLLogger.log("[SDLContext] udp running!!!!") - try await udpHole.waitClose() + try await self.startUDPHole() SDLLogger.log("[SDLContext] udp closed!!!!") } await self.supervisor.addWorker(name: "quicClient") { + SDLLogger.log("[SDLContext] superClient running!!!!") try await self.startQUICClient() + SDLLogger.log("[SDLContext] superClient closed!!!!") } // await self.supervisor.addWorker(name: "udpHoleV6") { @@ -179,6 +184,10 @@ actor SDLContextActor { SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") try await withThrowingTaskGroup { group in + defer { + group.cancelAll() + } + // 创建一个简单的异步状态等待机制(可以用一个 Actor 或者 AsyncStream 模拟) let (readyStream, readyContinuation) = AsyncStream.makeStream() @@ -339,31 +348,55 @@ actor SDLContextActor { } } - private func startUDPHole() async throws -> SDLUDPHole { - self.udpHoleWorkers?.forEach {$0.cancel()} - self.udpHoleWorkers = nil - + private func startUDPHole() async throws { // 启动udp服务器 let udpHole = try SDLUDPHole() let localAddress = try udpHole.start() SDLLogger.log("[SDLContext] udpHole started, on address: \(localAddress)") - - // 处理消息流 - let messageStream = udpHole.messageStream - let messageTask = Task.detached { - await self.consumeUDPHoleMessages(stream: messageStream, localAddress: localAddress, source: .v4) - } - self.udpHole = udpHole self.udpHoleLocalAddress = localAddress - self.udpHoleWorkers = [messageTask] - // 开始探测nat的类型 - Task { - await self.probeNatType() + defer { + self.udpHole = nil + self.udpHoleLocalAddress = nil } - return udpHole + try await withThrowingTaskGroup { group in + defer { + group.cancelAll() + } + + group.addTask { + for await (remoteAddress, message) in udpHole.messageStream { + try Task.checkCancellation() + + switch message.inboundMessage { + case .control(let controlMessage): + await self.handleHoleControlMessage(controlMessage, localAddress: localAddress, remoteAddress: remoteAddress, source: .v4) + case .data(let data): + try? await self.handleHoleData(data: data) + } + } + } + + group.addTask { + for await event in udpHole.eventStream { + try Task.checkCancellation() + switch event { + case .ready: + // 开始探测nat的类型 + Task { + await self.probeNatType() + } + SDLLogger.log("[SDLContext] udpHole ready") + case .closed, .errorCaught: + throw SDLContextError.udpHoleClosed + } + } + } + + try await group.next() + } } private func startUDPHoleV6() async throws -> SDLUDPHoleV6 { @@ -402,9 +435,6 @@ actor SDLContextActor { self.flowSessionManager.clear() - self.udpHoleWorkers?.forEach { $0.cancel() } - self.udpHoleWorkers = nil - self.udpHole?.stop() self.udpHole = nil self.udpHoleLocalAddress = nil diff --git a/Tun/Punchnet/UDPHole/SDLUDPHole.swift b/Tun/Punchnet/UDPHole/SDLUDPHole.swift index 30c7341..57a72d5 100644 --- a/Tun/Punchnet/UDPHole/SDLUDPHole.swift +++ b/Tun/Punchnet/UDPHole/SDLUDPHole.swift @@ -9,31 +9,40 @@ import NIOCore import NIOPosix import SwiftProtobuf +enum SDLUDPHoleError: Error { + case invalidLocalAddress +} + // 处理和sn-server服务器之间的通讯 final class SDLUDPHole: ChannelInboundHandler { typealias InboundIn = AddressedEnvelope - private enum State: Equatable { - case idle + // 事件 + enum HoleEvent { case ready - case stopping - case stopped + case closed + case errorCaught } 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 { @@ -45,55 +54,20 @@ final class SDLUDPHole: ChannelInboundHandler { // 绑定到IPv4通配地址,只处理IPv4流量 let channel = try bootstrap.bind(host: "0.0.0.0", port: 0).wait() + guard let localAddress = channel.localAddress else { + throw SDLUDPHoleError.invalidLocalAddress + } + self.channel = channel - self.closeFuture = channel.closeFuture - self.state = .ready - precondition(channel.localAddress != nil, "UDP channel has no localAddress after bind") + eventContinuation.yield(.ready) - return channel.localAddress! - } - - func waitClose() async throws { - switch self.state { - case .idle: - SDLLogger.log("[SDLUDPHole] waitClose11", for: .debug) - return - case .ready, .stopping, .stopped: - guard let closeFuture = self.closeFuture else { - SDLLogger.log("[SDLUDPHole] waitClose22", for: .debug) - return - } - try await closeFuture.get() - SDLLogger.log("[SDLUDPHole] waitClose33", for: .debug) - } - } - - func stop() { - SDLLogger.log("[SDLUDPHole] waitClose stop", for: .debug) - 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 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 let remoteAddress = envelope.remoteAddress @@ -113,25 +87,19 @@ final class SDLUDPHole: ChannelInboundHandler { } func channelInactive(context: ChannelHandlerContext) { - self.finishMessageStream() - self.channel = nil - self.state = .stopped SDLLogger.log("[SDLUDPHole] channelInactive", for: .debug) + self.eventContinuation.yield(.closed) } func errorCaught(context: ChannelHandlerContext, error: any Error) { SDLLogger.log("[SDLUDPHole] channel error: \(error)", for: .debug) - self.finishMessageStream() - if self.state != .stopped { - self.state = .stopping - } context.close(promise: nil) - SDLLogger.log("[SDLUDPHole] errorCaught", for: .debug) + 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 } @@ -145,18 +113,12 @@ final class SDLUDPHole: ChannelInboundHandler { } } - private func finishMessageStream() { - guard !self.didFinishMessageStream else { - return - } - - self.didFinishMessageStream = true - self.messageContinuation.finish() - } - deinit { - SDLLogger.log("[SDLUDPHole] closeWait deinit", for: .debug) - self.stop() + SDLLogger.log("[SDLUDPHole] deinit", for: .debug) + self.messageContinuation.finish() + self.eventContinuation.finish() + self.channel = nil + try? self.group.syncShutdownGracefully() }