From 6ac8ebf398dddc0ea158bdaf4b6b61fff9858073 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Tue, 28 Apr 2026 20:44:22 +0800 Subject: [PATCH] fix quicClient --- Tun/Punchnet/Actors/SDLContextActor.swift | 130 +++++--------- Tun/Punchnet/Actors/SDLQuicClient.swift | 203 +++++++--------------- Tun/Punchnet/AsyncOneShot.swift | 2 +- 3 files changed, 107 insertions(+), 228 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 5163939..a0d40f6 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -80,7 +80,6 @@ actor SDLContextActor { private var dnsLocalWorker: Task? private var quicClient: SDLQUICClient? - private var quicWorker: Task? nonisolated private let puncherActor: SDLPuncherActor // 网络探测对象 @@ -258,94 +257,63 @@ actor SDLContextActor { private func startQUICClient() async throws { SDLLogger.log("[SDLContext] try start quicClient", for: .debug) - - self.quicWorker?.cancel() - await self.quicClient?.stop() - + // 启动monitor let quicClient = SDLQUICClient(host: self.config.serverHost, port: 443) self.quicClient = quicClient - - await quicClient.start() - - do { - try await quicClient.waitReady(timeout: .seconds(3)) - // 这里必须等待quic的协商完成 - try await Task.sleep(for: .seconds(0.3)) - SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") - - try await withTaskCancellationHandler { - try await withThrowingTaskGroup { group in - defer { - group.cancelAll() + quicClient.start() + + defer { + quicClient.stop() + } + + try await quicClient.waitReady(timeout: .seconds(3)) + // 这里必须等待quic的协商完成 + try await Task.sleep(for: .seconds(0.3)) + SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") + + try await withThrowingTaskGroup { group in + defer { + group.cancelAll() + } + + group.addTask { + for try await message in quicClient.messageStream() { + await self.handleQUICMessage(message: message) + } + } + + group.addTask { + let timerStream = SDLAsyncTimerStream() + timerStream.start(interval: .seconds(5)) + + for await _ in timerStream.stream { + if Task.isCancelled { + break } - - group.addTask { - for await message in await quicClient.messageStream { - if Task.isCancelled { - return - } - await self.handleQUICMessage(message: message) - } - if Task.isCancelled { - return - } - throw SDLQUICClientExit.transportClosed("messageStream finished") - } - - group.addTask { - let exit = await quicClient.run() - - switch exit { - case .normal: - return - case .cancelled: - if Task.isCancelled { - return - } - throw exit - case .transportClosed, .readFailed, .writeFailed: - throw exit - } - } - - group.addTask { - for await event in await quicClient.eventStream { - switch event { - case .failed(let error): - throw error - case .cancelled: - throw SDLQUICClientExit.cancelled - case .writeFailed(let error): - throw error - } - } - if Task.isCancelled { - return - } - throw SDLQUICClientExit.transportClosed("eventStream finished") - } - - do { - let _ = try await group.next() - await quicClient.stop() - } catch { - await quicClient.stop() + quicClient.send(type: .ping, data: Data()) + } + SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug) + } + + group.addTask { + for await event in quicClient.eventStream { + switch event { + case .failed(let error): + throw error + case .cancelled: + throw SDLQUICEvent.cancelled + case .writeFailed(let error): throw error } } - - } onCancel: { - Task { - await quicClient.stop() - } } - } catch { - await quicClient.stop() - throw error + + try await group.next() } + } - + private func handleQUICMessage(message: SDLQUICInboundMessage) async { switch message { case .welcome(let welcome): @@ -533,8 +501,6 @@ actor SDLContextActor { self.udpHoleV6 = nil self.udpHoleV6LocalAddress = nil - self.quicWorker?.cancel() - self.quicWorker = nil await self.quicClient?.stop() self.quicClient = nil @@ -996,7 +962,7 @@ extension SDLContextActor { if let registerSuperData = try? registerSuper.serializedData() { SDLLogger.log("[SDLContext] will send register super") - await self.quicClient?.send(type: .registerSuper, data: registerSuperData) + self.quicClient?.send(type: .registerSuper, data: registerSuperData) } } diff --git a/Tun/Punchnet/Actors/SDLQuicClient.swift b/Tun/Punchnet/Actors/SDLQuicClient.swift index b5184ee..92e30dc 100644 --- a/Tun/Punchnet/Actors/SDLQuicClient.swift +++ b/Tun/Punchnet/Actors/SDLQuicClient.swift @@ -26,57 +26,26 @@ enum SDLQUICEvent: Error { case writeFailed(Error) } -enum SDLQUICClientExit: Error, Sendable, CustomStringConvertible { - case normal - case cancelled - case transportClosed(String) - case readFailed(String) - case writeFailed(String) - - var description: String { - switch self { - case .normal: - return "normal" - case .cancelled: - return "cancelled" - case .transportClosed(let reason): - return "transportClosed(\(reason))" - case .readFailed(let reason): - return "readFailed(\(reason))" - case .writeFailed(let reason): - return "writeFailed(\(reason))" - } - } -} - -actor SDLQUICClient { +final class SDLQUICClient { private let allocator = ByteBufferAllocator() - // 单个包最大64K - private let maxPacketSize: Int // 最大缓冲区区为2M private let maxBufferSize: Int + + private static let maxPacketSize: Int = 64 * 1024 private let readyLatch = AsyncOneShot() - // 消息流 - public var messageStream: AsyncStream - private let messageCont: AsyncStream.Continuation - // 事件流 public var eventStream: AsyncStream private let eventCont: AsyncStream.Continuation - private var didFinishStreams = false - private let connection: NWConnection private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 专用队列保证线程安全 - init(host: String, port: UInt16, maxPacketSize: Int = 64 * 1024, maxBufferSize: Int = 2 * 1024 * 1024) { + init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) { let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"]) self.maxBufferSize = maxBufferSize - self.maxPacketSize = maxPacketSize - (self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self) (self.eventStream, self.eventCont) = AsyncStream.makeStream(of: SDLQUICEvent.self) // 这里设置证书的校验逻辑 @@ -101,7 +70,7 @@ actor SDLQUICClient { } connection.start(queue: self.queue) } - + private func handleConnectionStateUpdate(_ state: NWConnection.State) async { SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) switch state { @@ -109,10 +78,10 @@ actor SDLQUICClient { await self.readyLatch.succeed(()) case .failed(let error): await self.readyLatch.fail(error) - self.emitEvent(.failed(error)) + self.eventCont.yield(.failed(error)) case .cancelled: await self.readyLatch.fail(SDLQUICError.connectionCancelled) - self.emitEvent(.cancelled) + self.eventCont.yield(.cancelled) default: () } @@ -122,24 +91,6 @@ actor SDLQUICClient { try await self.readyLatch.wait(timeout: timeout, timeoutError: SDLQUICError.timeout) } - func run() async -> SDLQUICClientExit { - await withTaskGroup(of: SDLQUICClientExit.self) { group in - group.addTask { - await self.readLoop() - } - - group.addTask { - await self.heartbeatLoop() - } - - let exit = await group.next() ?? .normal - group.cancelAll() - self.finishStreams() - - return exit - } - } - func send(type: SDLPacketType, data: Data) { var len = UInt16(data.count + 1).bigEndian @@ -150,139 +101,101 @@ actor SDLQUICClient { connection.send(content: packet, completion: .contentProcessed { [weak self] error in if let error { SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug) - Task { - await self?.emitEvent(.writeFailed(error)) - } + self?.eventCont.yield(.writeFailed(error)) } }) } - private func heartbeatLoop() async -> SDLQUICClientExit { - let timerStream = SDLAsyncTimerStream() - timerStream.start(interval: .seconds(5)) - - for await _ in timerStream.stream { - if Task.isCancelled { - break - } - self.send(type: .ping, data: Data()) - } - - SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug) - return .cancelled - } - - func stop() async { + func stop() { self.connection.cancel() - await self.readyLatch.fail(SDLQUICError.connectionCancelled) - self.finishStreams() - } - - private func emitEvent(_ event: SDLQUICEvent) { - guard !self.didFinishStreams else { - return + Task { + await self.readyLatch.fail(SDLQUICError.connectionCancelled) } - - self.eventCont.yield(event) - } - - private func finishStreams() { - guard !self.didFinishStreams else { - return - } - - self.didFinishStreams = true - self.messageCont.finish() self.eventCont.finish() } - + } // --MARK: Reader extension SDLQUICClient { - private func readLoop() async -> SDLQUICClientExit { - var buffer = allocator.buffer(capacity: self.maxBufferSize) - let threshold = self.maxBufferSize / 10 * 6 - - defer { - self.messageCont.finish() - } - - do { - while !Task.isCancelled { - let (isComplete, data) = try await self.readOnce() - if let data, !data.isEmpty { - buffer.writeBytes(data) - let frames = try parseFrames(buffer: &buffer) - if buffer.readerIndex > threshold { - buffer.discardReadBytes() + + func messageStream() -> AsyncThrowingStream { + return AsyncThrowingStream { continuation in + var buffer = allocator.buffer(capacity: self.maxBufferSize) + let threshold = self.maxBufferSize / 10 * 6 + + func readOnce() { + self.connection.receive(minimumIncompleteLength: 1, maximumLength: Self.maxPacketSize) { data, _, isComplete, error in + if let error { + continuation.finish(throwing: error) + return } - - for frame in frames { - if let message = decode(frame: frame) { - self.messageCont.yield(message) + + do { + if let data, !data.isEmpty { + buffer.writeBytes(data) + let frames = try Self.parseFrames(buffer: &buffer) + if buffer.readerIndex > threshold { + buffer.discardReadBytes() + } + + for frame in frames { + if let message = Self.decode(frame: frame) { + continuation.yield(message) + } + } } + } catch let err { + continuation.finish(throwing: err) + return + } + + if isComplete { + continuation.finish() + } else { + readOnce() } - } - - if isComplete { - return .transportClosed("receive complete") } } - return .cancelled - } catch is CancellationError { - return .cancelled - } catch { - return .readFailed("\(error)") + readOnce() } } - + // 尝试解析数据 - private func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] { + private static func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] { guard buffer.readableBytes >= 2 else { return [] } - + var frames: [ByteBuffer] = [] while true { guard let len = buffer.getInteger(at: buffer.readerIndex, endianness: .big, as: UInt16.self) else { break } - - if len > self.maxPacketSize { + + if len > Self.maxPacketSize { throw SDLQUICError.packetTooLarge } - + guard buffer.readableBytes >= len + 2 else { break } - + buffer.moveReaderIndex(forwardBy: 2) if let buf = buffer.readSlice(length: Int(len)) { frames.append(buf) } } - + return frames } - - // 读取一次数据 - private func readOnce() async throws -> (Bool, Data?) { - return try await withCheckedThrowingContinuation { cont in - self.connection.receive(minimumIncompleteLength: 1, maximumLength: maxPacketSize) { data, _, isComplete, error in - if let error { - cont.resume(throwing: error) - return - } - cont.resume(returning: (isComplete, data)) - } - } - } + } // --MARK: 编解码器 extension SDLQUICClient { - private func decode(frame: ByteBuffer) -> SDLQUICInboundMessage? { + + private static func decode(frame: ByteBuffer) -> SDLQUICInboundMessage? { var buffer = frame guard let type = buffer.readInteger(as: UInt8.self), let packetType = SDLPacketType(rawValue: type) else { diff --git a/Tun/Punchnet/AsyncOneShot.swift b/Tun/Punchnet/AsyncOneShot.swift index a50d140..a3f7d72 100644 --- a/Tun/Punchnet/AsyncOneShot.swift +++ b/Tun/Punchnet/AsyncOneShot.swift @@ -71,7 +71,7 @@ actor AsyncOneShot { do { try await Task.sleep(for: timeout) if !Task.isCancelled { - await self.cancelWaiter(id: id, throwing: timeoutError) + self.cancelWaiter(id: id, throwing: timeoutError) } } catch { return