From a2901a0f8fb268cd26d1dbd761f00c17029b7220 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Tue, 5 May 2026 17:38:30 +0800 Subject: [PATCH] =?UTF-8?q?=E8=B0=83=E6=95=B4=E7=94=9F=E5=91=BD=E5=91=A8?= =?UTF-8?q?=E6=9C=9F=E7=9A=84=E7=AE=A1=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Tun/Punchnet/Actors/SDLContextActor.swift | 84 +++++++------- Tun/Punchnet/Actors/SDLNATProberActor.swift | 8 +- Tun/Punchnet/Actors/SDLPuncherActor.swift | 8 +- Tun/Punchnet/UDPHole/SDLUDPHole.swift | 115 ++++++++++++++------ 4 files changed, 131 insertions(+), 84 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 4708525..e35ade0 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -166,6 +166,10 @@ actor SDLContextActor { self.quicClient = quicClient await quicClient.start() + defer { + self.quicClient = nil + } + // 这里必须等待quic的协商完成 try await Task.sleep(for: .seconds(0.5)) SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") @@ -284,6 +288,11 @@ actor SDLContextActor { let dnsClient = DNSCloudClient(host: self.config.serverHost, port: 15353) self.dnsClient = dnsClient dnsClient.start() + + defer { + self.dnsClient = nil + } + do { try await withTaskCancellationHandler { for try await packet in dnsClient.packetFlow { @@ -319,6 +328,10 @@ actor SDLContextActor { SDLLogger.log("[SDLContext] dnsLocalClient started") self.dnsLocalClient = dnsLocalClient + defer { + self.dnsLocalClient = nil + } + do { try await withTaskCancellationHandler { // 处理事件流 @@ -352,24 +365,24 @@ actor SDLContextActor { private func startUDPHole() async throws { // 启动udp服务器 let udpHole = try SDLUDPHole() - let localAddress = try udpHole.start() + let localAddress = try await udpHole.start() SDLLogger.log("[SDLContext] udpHole started, on address: \(localAddress)") self.udpHole = udpHole self.udpHoleLocalAddress = localAddress defer { - self.udpHole?.stop() self.udpHole = nil - self.udpHoleLocalAddress = nil } - try await withThrowingTaskGroup { group in - defer { - group.cancelAll() - } - - group.addTask { - for await (remoteAddress, message) in udpHole.messageStream { + // 开始探测nat的类型 + Task { + await self.probeNatType() + } + SDLLogger.log("[SDLContext] udpHole ready") + + do { + try await withTaskCancellationHandler { + for try await (remoteAddress, message) in await udpHole.messageStream() { try Task.checkCancellation() switch message.inboundMessage { @@ -388,25 +401,14 @@ actor SDLContextActor { 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 - } + } onCancel: { + Task { + await udpHole.stop() } } - - try await group.next() + } catch let err { + await udpHole.stop() + throw err } } @@ -564,7 +566,7 @@ actor SDLContextActor { } } - private func sendStunRequest(probeReply: SDLV6AssistProbeReply?) { + private func sendStunRequest(probeReply: SDLV6AssistProbeReply?) async { guard let sessionToken = self.sessionToken else { return } @@ -582,7 +584,7 @@ actor SDLContextActor { } if let stunData = try? stunRequest.serializedData() { - self.sendSuperPacket(type: .stunRequest, data: stunData) + await self.sendSuperPacket(type: .stunRequest, data: stunData) } } @@ -661,23 +663,23 @@ actor SDLContextActor { } // 发送给super/stun节点的数据 - private func sendSuperPacket(type: SDLPacketType, data: Data) { - self.sendPacket(type: type, data: data, remoteAddress: self.config.stunSocketAddress) + private func sendSuperPacket(type: SDLPacketType, data: Data) async { + await self.sendPacket(type: type, data: data, remoteAddress: self.config.stunSocketAddress) } // 发送给peer的数据 - private func sendPeerPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { - self.sendPacket(type: type, data: data, remoteAddress: remoteAddress) + private func sendPeerPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { + await self.sendPacket(type: type, data: data, remoteAddress: remoteAddress) } - private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { + private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { switch remoteAddress { case .v4: guard let udpHole = self.udpHole else { SDLLogger.log("[SDLContext] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) return } - udpHole.send(type: type, data: data, remoteAddress: remoteAddress) + await udpHole.send(type: type, data: data, remoteAddress: remoteAddress) case .v6: guard let udpHoleV6 = self.udpHoleV6 else { SDLLogger.log("[SDLContext] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) @@ -821,8 +823,8 @@ extension SDLContextActor { case .removeSession(let dstMac): await self.sessionManager.removeSession(dstMac: dstMac) case .sendRegister(let registerData, let remoteAddresses): - remoteAddresses.forEach { remoteAddress in - self.sendPeerPacket(type: .register, data: registerData, remoteAddress: remoteAddress) + for remoteAddress in remoteAddresses { + await self.sendPeerPacket(type: .register, data: registerData, remoteAddress: remoteAddress) } case .shutdown(let message): self.publishTunnelEvent(message: message) @@ -870,7 +872,7 @@ extension SDLContextActor { registerAck.srcMac = networkAddr.mac registerAck.dstMac = register.srcMac - self.sendPeerPacket(type: .registerAck, data: try registerAck.serializedData(), remoteAddress: remoteAddress) + await self.sendPeerPacket(type: .registerAck, data: try registerAck.serializedData(), remoteAddress: remoteAddress) // 这里需要建立到来源的会话, 在复杂网络下,通过super-node查询到的nat地址不一定靠谱,需要通过udp包的来源地址作为nat地址 if let session = Session(dstMac: register.srcMac, natAddress: remoteAddress, addressType: source.convertAddressType()) { await self.sessionManager.addSession(session: session) @@ -1024,15 +1026,15 @@ extension SDLContextActor { switch plan { case .superNode(let payload): // 通过super_node进行转发 - self.sendSuperPacket(type: .data, data: payload) + await self.sendSuperPacket(type: .data, data: payload) case .peer(let payload, let session): // 通过session发送到对端 SDLLogger.log("[SDLContext] step 5 send packet by session: \(session)", for: .trace) - self.sendPeerPacket(type: .data, data: payload, remoteAddress: session.natAddress) + await self.sendPeerPacket(type: .data, data: payload, remoteAddress: session.natAddress) self.flowTracer.inc(num: payload.count, type: .p2p) case .superNodeAndPunch(let payload, let request): // 通过super_node进行转发 - self.sendSuperPacket(type: .data, data: payload) + await self.sendSuperPacket(type: .data, data: payload) SDLLogger.log("[SDLContext] step 5 send packet by super: \(self.config.stunSocketAddress)", for: .trace) // 流量统计 self.flowTracer.inc(num: payload.count, type: .forward) diff --git a/Tun/Punchnet/Actors/SDLNATProberActor.swift b/Tun/Punchnet/Actors/SDLNATProberActor.swift index 93c2c99..9ba707c 100644 --- a/Tun/Punchnet/Actors/SDLNATProberActor.swift +++ b/Tun/Punchnet/Actors/SDLNATProberActor.swift @@ -158,10 +158,10 @@ actor SDLNATProberActor { // MARK: - Internal helpers private func sendProbe(using udpHole: SDLUDPHole, cookie: UInt32) async { - udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 1, attr: .none), remoteAddress: addressArray[0][0]) - udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 2, attr: .none), remoteAddress: addressArray[1][1]) - udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 3, attr: .peer), remoteAddress: addressArray[0][0]) - udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 4, attr: .port), remoteAddress: addressArray[0][0]) + await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 1, attr: .none), remoteAddress: addressArray[0][0]) + await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 2, attr: .none), remoteAddress: addressArray[1][1]) + await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 3, attr: .peer), remoteAddress: addressArray[0][0]) + await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 4, attr: .port), remoteAddress: addressArray[0][0]) } private func makeProbePacket(cookieId: UInt32, step: UInt32, attr: SDLProbeAttr) -> Data { diff --git a/Tun/Punchnet/Actors/SDLPuncherActor.swift b/Tun/Punchnet/Actors/SDLPuncherActor.swift index 2ba976c..1a8ea83 100644 --- a/Tun/Punchnet/Actors/SDLPuncherActor.swift +++ b/Tun/Punchnet/Actors/SDLPuncherActor.swift @@ -127,7 +127,7 @@ actor SDLPuncherActor { if peerInfo.hasV4Info { if let remoteAddress = try? await peerInfo.v4Info.socketAddress() { SDLLogger.log("[SDLContext] hole sock address: \(remoteAddress)", for: .debug) - self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress) + await self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress) } else { SDLLogger.log("[SDLPuncherActor] failed to resolve peerInfo.v4Info", for: .debug) } @@ -136,7 +136,7 @@ actor SDLPuncherActor { if peerInfo.hasV6Info { if let remoteAddress = try? await peerInfo.v6Info.socketAddress() { SDLLogger.log("[SDLContext] hole sock address v6: \(remoteAddress)", for: .debug) - self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress) + await self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress) } else { SDLLogger.log("[SDLPuncherActor] failed to resolve peerInfo.v6Info", for: .debug) } @@ -156,14 +156,14 @@ actor SDLPuncherActor { } } - private func sendRegister(using udpHole: SDLUDPHole?, udpHoleV6: SDLUDPHoleV6?, registerData: Data, remoteAddress: SocketAddress) { + private func sendRegister(using udpHole: SDLUDPHole?, udpHoleV6: SDLUDPHoleV6?, registerData: Data, remoteAddress: SocketAddress) async { switch remoteAddress { case .v4: guard let udpHole else { SDLLogger.log("[SDLPuncherActor] udpHole is nil when v4 peerInfo arrived", for: .debug) return } - udpHole.send(type: .register, data: registerData, remoteAddress: remoteAddress) + await udpHole.send(type: .register, data: registerData, remoteAddress: remoteAddress) case .v6: guard let udpHoleV6 else { SDLLogger.log("[SDLPuncherActor] udpHoleV6 is nil when v6 peerInfo arrived", for: .debug) diff --git a/Tun/Punchnet/UDPHole/SDLUDPHole.swift b/Tun/Punchnet/UDPHole/SDLUDPHole.swift index ea0f895..93718f1 100644 --- a/Tun/Punchnet/UDPHole/SDLUDPHole.swift +++ b/Tun/Punchnet/UDPHole/SDLUDPHole.swift @@ -11,40 +11,73 @@ import SwiftProtobuf enum SDLUDPHoleError: Error { case invalidLocalAddress + case closed + case errorCaught + case sendFaied(Error) +} + +actor SDLUDPHole { + enum State { + case idle + case running + case stopped + } + + private var state: State = .idle + + private let udpHoleHandler: SDLUDPHoleHandler + + init() throws { + self.udpHoleHandler = try SDLUDPHoleHandler() + } + + func start() throws -> SocketAddress { + let localAddress = try self.udpHoleHandler.start() + self.state = .running + + return localAddress + } + + func messageStream() -> AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error> { + return self.udpHoleHandler.messageStream + } + + func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { + guard self.state == .running else { + return + } + + self.udpHoleHandler.send(type: type, data: data, remoteAddress: remoteAddress) + } + + func stop() { + guard self.state != .stopped else { + return + } + self.state = .stopped + self.udpHoleHandler.stop() + } + } // 处理和sn-server服务器之间的通讯 -final class SDLUDPHole: ChannelInboundHandler { +private final class SDLUDPHoleHandler: ChannelInboundHandler { typealias InboundIn = AddressedEnvelope - // 事件 - enum HoleEvent { - case ready - case closed - case errorCaught - } - - private var isStopped: Bool = false - private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) private var channel: Channel? - public let messageStream: AsyncStream<(SocketAddress, SDLHoleMessage)> - private let messageContinuation: AsyncStream<(SocketAddress, SDLHoleMessage)>.Continuation + private let locker = NSLock() + + public let messageStream: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error> + private let messageContinuation: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error>.Continuation + private var isMessageContinuationFinished: Bool = false - // 事件相关逻辑 - public let eventStream: AsyncStream - private let eventContinuation: AsyncStream.Continuation - // 启动函数 init() throws { - let (stream, continuation) = AsyncStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048)) + let (stream, continuation) = AsyncThrowingStream.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 { @@ -61,7 +94,6 @@ final class SDLUDPHole: ChannelInboundHandler { } self.channel = channel - eventContinuation.yield(.ready) return localAddress } @@ -89,13 +121,13 @@ final class SDLUDPHole: ChannelInboundHandler { func channelInactive(context: ChannelHandlerContext) { SDLLogger.log("[SDLUDPHole] channelInactive", for: .debug) - self.eventContinuation.yield(.closed) + self.finishMessageContinuationIfNeed(throwing: .closed) } func errorCaught(context: ChannelHandlerContext, error: any Error) { SDLLogger.log("[SDLUDPHole] channel error: \(error)", for: .debug) context.close(promise: nil) - self.eventContinuation.yield(.errorCaught) + self.finishMessageContinuationIfNeed(throwing: .errorCaught) } // MARK: 处理写入逻辑 @@ -109,24 +141,37 @@ final class SDLUDPHole: ChannelInboundHandler { buffer.writeBytes(data) let envelope = AddressedEnvelope(remoteAddress: remoteAddress, data: buffer) - _ = channel.eventLoop.submit { - channel.writeAndFlush(envelope, promise: nil) + + let promise = channel.eventLoop.makePromise(of: Void.self) + channel.eventLoop.execute { + channel.writeAndFlush(envelope, promise: promise) + } + + promise.futureResult.whenFailure { err in + self.finishMessageContinuationIfNeed(throwing: .sendFaied(err)) } } func stop() { - guard !self.isStopped else { - return - } - - self.isStopped = true - SDLLogger.log("[SDLUDPHole] stop", for: .debug) - self.messageContinuation.finish() - self.eventContinuation.finish() + self.finishMessageContinuationIfNeed(throwing: nil) + try? self.channel?.close().wait() self.channel = nil - try? self.group.syncShutdownGracefully() } + private func finishMessageContinuationIfNeed(throwing error: SDLUDPHoleError?) { + locker.lock() + defer { + locker.unlock() + } + + guard !self.isMessageContinuationFinished else { + return + } + + self.isMessageContinuationFinished = true + self.messageContinuation.finish(throwing: error) + } + }