From e07c81567a067886409f3551e4c7c50fb639d263 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Fri, 22 May 2026 17:04:05 +0800 Subject: [PATCH] fix udpHole --- Tun/UDPHole/SDLUDPHole.swift | 125 +-------------------------- Tun/UDPHole/SDLUDPHoleError.swift | 14 +++ Tun/UDPHole/SDLUDPHoleHandler.swift | 128 ++++++++++++++++++++++++++++ Tun/UDPHole/SDLUDPHoleService.swift | 10 +-- 4 files changed, 146 insertions(+), 131 deletions(-) create mode 100644 Tun/UDPHole/SDLUDPHoleError.swift create mode 100644 Tun/UDPHole/SDLUDPHoleHandler.swift diff --git a/Tun/UDPHole/SDLUDPHole.swift b/Tun/UDPHole/SDLUDPHole.swift index 86c674e..f160d71 100644 --- a/Tun/UDPHole/SDLUDPHole.swift +++ b/Tun/UDPHole/SDLUDPHole.swift @@ -9,13 +9,6 @@ import NIOCore import NIOPosix import SwiftProtobuf -enum SDLUDPHoleError: Error { - case invalidLocalAddress - case closed - case errorCaught - case sendFaied(Error) -} - actor SDLUDPHole { enum State { case idle @@ -37,7 +30,7 @@ actor SDLUDPHole { return localAddress } - func messageStream() -> AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error> { + func messageStream() -> AsyncThrowingStream { return self.udpHoleHandler.messageStream } @@ -63,119 +56,3 @@ actor SDLUDPHole { } } -// 处理和sn-server服务器之间的通讯 -private final class SDLUDPHoleHandler: ChannelInboundHandler { - typealias InboundIn = AddressedEnvelope - - private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) - private var channel: Channel? - - private let locker = NSLock() - - public let messageStream: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error> - private let messageContinuation: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error>.Continuation - private var isMessageContinuationFinished: Bool = false - - // 启动函数 - init() throws { - let (stream, continuation) = AsyncThrowingStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048)) - self.messageStream = stream - self.messageContinuation = continuation - } - - func start() throws -> SocketAddress { - let bootstrap = DatagramBootstrap(group: group) - .channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) - .channelInitializer { channel in - channel.pipeline.addHandler(self) - } - - // 绑定到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 - - return localAddress - } - - // --MARK: ChannelInboundHandler delegate - - func channelRead(context: ChannelHandlerContext, data: NIOAny) { - let envelope = unwrapInboundIn(data) - var buffer = envelope.data - let remoteAddress = envelope.remoteAddress - - do { - if let message = try SDLHoleMessage.decode(buffer: &buffer) { - self.messageContinuation.yield((remoteAddress, message)) - } else { - SDLLogger.log("[SDLUDPHole] decode message, get null", for: .debug) - } - } catch let err { - SDLLogger.log("[SDLUDPHole] decode message, get error: \(err)", for: .debug) - } - } - - func channelInactive(context: ChannelHandlerContext) { - self.finishMessageContinuationIfNeed(throwing: .closed) - } - - func errorCaught(context: ChannelHandlerContext, error: any Error) { - context.close(promise: nil) - self.finishMessageContinuationIfNeed(throwing: .errorCaught) - } - - // MARK: 处理写入逻辑 - func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { - guard let channel = self.channel else { - return - } - - var buffer = channel.allocator.buffer(capacity: data.count + 1) - buffer.writeBytes([type.rawValue]) - buffer.writeBytes(data) - - let envelope = AddressedEnvelope(remoteAddress: remoteAddress, data: buffer) - - let promise = channel.eventLoop.makePromise(of: Void.self) - channel.eventLoop.execute { - channel.writeAndFlush(envelope, promise: promise) - } - - promise.futureResult.whenFailure { [weak self] err in - self?.finishMessageContinuationIfNeed(throwing: .sendFaied(err)) - } - } - - func stop() { - self.finishMessageContinuationIfNeed(throwing: nil) - let channel = self.channel - self.channel = nil - try? channel?.close().wait() - try? self.group.syncShutdownGracefully() - - SDLLogger.log("[SDLUDPHole] stopped", for: .debug) - } - - private func finishMessageContinuationIfNeed(throwing error: SDLUDPHoleError?) { - locker.lock() - defer { - locker.unlock() - } - - guard !self.isMessageContinuationFinished else { - return - } - - self.isMessageContinuationFinished = true - self.messageContinuation.finish(throwing: error) - } - - deinit { - SDLLogger.log("[SDLUDPHole] deinit", for: .debug) - } - -} diff --git a/Tun/UDPHole/SDLUDPHoleError.swift b/Tun/UDPHole/SDLUDPHoleError.swift new file mode 100644 index 0000000..83f1479 --- /dev/null +++ b/Tun/UDPHole/SDLUDPHoleError.swift @@ -0,0 +1,14 @@ +// +// SDLUDPHoleError.swift +// punchnet +// +// Created by 安礼成 on 2026/5/22. +// +import Foundation + +enum SDLUDPHoleError: Error { + case invalidLocalAddress + case closed + case errorCaught + case sendFaied(Error) +} diff --git a/Tun/UDPHole/SDLUDPHoleHandler.swift b/Tun/UDPHole/SDLUDPHoleHandler.swift new file mode 100644 index 0000000..c3ce320 --- /dev/null +++ b/Tun/UDPHole/SDLUDPHoleHandler.swift @@ -0,0 +1,128 @@ +// +// SDLUDPHoleHandler.swift +// punchnet +// +// Created by 安礼成 on 2026/5/22. +// +import Foundation +import NIOCore +import NIOPosix + +// 处理和sn-server服务器之间的通讯 +final class SDLUDPHoleHandler: ChannelInboundHandler { + typealias InboundIn = AddressedEnvelope + + struct SDLHoleDatagram { + let remoteAddress: SocketAddress + let message: SDLHoleMessage + } + + private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) + private var channel: Channel? + + public let messageStream: AsyncThrowingStream + private let messageContinuation: AsyncThrowingStream.Continuation + + private let locker = NSLock() + private var isStopped: Bool = false + + // 启动函数 + init() throws { + let (stream, continuation) = AsyncThrowingStream.makeStream(of: SDLHoleDatagram.self, bufferingPolicy: .bufferingNewest(2048)) + self.messageStream = stream + self.messageContinuation = continuation + } + + func start() throws -> SocketAddress { + let bootstrap = DatagramBootstrap(group: group) + .channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) + .channelInitializer { channel in + channel.pipeline.addHandler(self) + } + + // 绑定到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 + + return localAddress + } + + // --MARK: ChannelInboundHandler delegate + + func channelRead(context: ChannelHandlerContext, data: NIOAny) { + let envelope = unwrapInboundIn(data) + var buffer = envelope.data + let remoteAddress = envelope.remoteAddress + + do { + if let message = try SDLHoleMessage.decode(buffer: &buffer) { + self.messageContinuation.yield(SDLHoleDatagram(remoteAddress: remoteAddress, message: message)) + } else { + SDLLogger.log("[SDLUDPHole] decode message, get null", for: .debug) + } + } catch let err { + SDLLogger.log("[SDLUDPHole] decode message, get error: \(err)", for: .debug) + self.messageContinuation.finish(throwing: err) + } + } + + func channelInactive(context: ChannelHandlerContext) { + self.messageContinuation.finish(throwing: SDLUDPHoleError.closed) + } + + func errorCaught(context: ChannelHandlerContext, error: any Error) { + context.close(promise: nil) + self.messageContinuation.finish(throwing: SDLUDPHoleError.errorCaught) + } + + // MARK: 处理写入逻辑 + func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { + guard let channel = self.channel else { + return + } + + var buffer = channel.allocator.buffer(capacity: data.count + 1) + buffer.writeBytes([type.rawValue]) + buffer.writeBytes(data) + + let envelope = AddressedEnvelope(remoteAddress: remoteAddress, data: buffer) + + let promise = channel.eventLoop.makePromise(of: Void.self) + channel.eventLoop.execute { + channel.writeAndFlush(envelope, promise: promise) + } + + promise.futureResult.whenFailure { [weak self] err in + self?.messageContinuation.finish(throwing: SDLUDPHoleError.sendFaied(err)) + } + } + + func stop() { + locker.lock() + defer { + locker.unlock() + } + + guard !self.isStopped else { + return + } + self.isStopped = true + + self.messageContinuation.finish() + let channel = self.channel + self.channel = nil + try? channel?.close().wait() + try? self.group.syncShutdownGracefully() + + SDLLogger.log("[SDLUDPHole] stopped", for: .debug) + } + + deinit { + SDLLogger.log("[SDLUDPHole] deinit", for: .debug) + } + +} diff --git a/Tun/UDPHole/SDLUDPHoleService.swift b/Tun/UDPHole/SDLUDPHoleService.swift index 9f13891..a1db8fd 100644 --- a/Tun/UDPHole/SDLUDPHoleService.swift +++ b/Tun/UDPHole/SDLUDPHoleService.swift @@ -38,11 +38,7 @@ actor SDLUDPHoleService { private var udpHoleV6: SDLUDPHoleV6? private var udpHoleV6MonitorTask: Task? - init( - proberActor: SDLNATProberActor, - onEvent: @escaping EventHandler, - onData: @escaping DataHandler - ) { + init(proberActor: SDLNATProberActor, onEvent: @escaping EventHandler, onData: @escaping DataHandler) { self.proberActor = proberActor self.onEvent = onEvent self.onData = onData @@ -140,9 +136,9 @@ actor SDLUDPHoleService { do { try await withTaskCancellationHandler { - for try await (remoteAddress, message) in await udpHole.messageStream() { + for try await datagram in await udpHole.messageStream() { try Task.checkCancellation() - try await self.handleV4Message(remoteAddress: remoteAddress, message: message) + try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message) } } onCancel: { Task {