From a70419e3e2730a45c3bd9c90146914e335ce333e Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Tue, 28 Apr 2026 14:18:08 +0800 Subject: [PATCH] fix quicClient --- Tun/Punchnet/Actors/SDLContextActor.swift | 10 +- Tun/Punchnet/Actors/SDLQuicClient.swift | 201 ++++++++++++++++------ 2 files changed, 161 insertions(+), 50 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 1e04c5d..1025688 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -168,8 +168,14 @@ actor SDLContextActor { SDLLogger.log("[SDLContext] try start quicClient", for: .debug) let quicClient = try await self.startQUICClient() SDLLogger.log("[SDLContext] quicClient running!!!!") - await quicClient.waitClose() - SDLLogger.log("[SDLContext] quicClient closed!!!!") + let exit = await quicClient.run() + SDLLogger.log("[SDLContext] quicClient closed: \(exit)") + switch exit { + case .normal, .cancelled: + return + case .transportClosed, .readFailed, .writeFailed: + throw exit + } } await self.supervisor.addWorker(name: "udpHole") { diff --git a/Tun/Punchnet/Actors/SDLQuicClient.swift b/Tun/Punchnet/Actors/SDLQuicClient.swift index 2fa951b..52c5819 100644 --- a/Tun/Punchnet/Actors/SDLQuicClient.swift +++ b/Tun/Punchnet/Actors/SDLQuicClient.swift @@ -20,6 +20,58 @@ enum SDLQUICError: Error { case packetTooLarge } +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))" + } + } +} + +private actor SDLQUICCloseWait { + private var exit: SDLQUICClientExit? + private var waiters: [CheckedContinuation] = [] + + func wait() async -> SDLQUICClientExit { + if let exit { + return exit + } + + return await withCheckedContinuation { continuation in + waiters.append(continuation) + } + } + + func close(_ exit: SDLQUICClientExit) { + guard self.exit == nil else { + return + } + + self.exit = exit + let waiters = self.waiters + self.waiters.removeAll() + + for waiter in waiters { + waiter.resume(returning: exit) + } + } +} + final class SDLQUICClient { private let allocator = ByteBufferAllocator() // 单个包最大64K @@ -35,7 +87,7 @@ final class SDLQUICClient { private let connection: NWConnection private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 专用队列保证线程安全 - private let (closeStream, closeCont) = AsyncStream.makeStream(of: Void.self) + private let closeWait = SDLQUICCloseWait() private let (readyStream, readyCont) = AsyncStream.makeStream(of: Void.self) init(host: String, port: UInt16, maxPacketSize: Int = 64 * 1024, maxBufferSize: Int = 2 * 1024 * 1024) { @@ -66,60 +118,50 @@ final class SDLQUICClient { case .ready: self.readyCont.yield() self.readyCont.finish() - case .failed(_), .cancelled: - self.closeCont.yield() - self.closeCont.finish() + case .failed(let error): + self.readyCont.finish() + Task { + await self.close(.transportClosed("failed: \(error)")) + } + case .cancelled: + self.readyCont.finish() + Task { + await self.close(.cancelled) + } default: () } } connection.start(queue: self.queue) - - // 启动数据读取任务 - self.readTask = Task { - var buffer = allocator.buffer(capacity: self.maxBufferSize) - let threshold = self.maxBufferSize / 10 * 6 - 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() - } - - for frame in frames { - if let message = decode(frame: frame) { - self.messageCont.yield(message) - } - } - } - - if isComplete { - break - } + } + + func run() async -> SDLQUICClientExit { + await withTaskCancellationHandler { + await withTaskGroup(of: SDLQUICClientExit.self) { group in + group.addTask { + await self.readLoop() } + group.addTask { + await self.heartbeatLoop() + } + group.addTask { + await self.waitClose() + } + + let exit = await group.next() ?? .normal + group.cancelAll() + self.connection.cancel() self.messageCont.finish() - } catch { + await self.close(exit) + return exit + } + } onCancel: { + Task { + await self.close(.cancelled) + self.connection.cancel() self.messageCont.finish() } } - - // 处理心跳逻辑 - self.pingTask = Task { - 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) - } } func send(type: SDLPacketType, data: Data) { @@ -132,22 +174,85 @@ final class SDLQUICClient { connection.send(content: packet, completion: .contentProcessed { error in if let error { SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug) + Task { + await self.close(.writeFailed("\(error)")) + } } }) } func waitReady() async throws { - for await _ in readyStream {} + for await _ in readyStream { + return + } + let exit = await closeWait.wait() + throw exit } - func waitClose() async { - for await _ in closeStream {} + func waitClose() async -> SDLQUICClientExit { + await closeWait.wait() } func stop() { self.connection.cancel() } + func close(_ exit: SDLQUICClientExit = .normal) async { + await closeWait.close(exit) + } + + 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() + } + + for frame in frames { + if let message = decode(frame: frame) { + self.messageCont.yield(message) + } + } + } + + if isComplete { + return .transportClosed("receive complete") + } + } + return .cancelled + } catch is CancellationError { + return .cancelled + } catch { + return .readFailed("\(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 + } + // 尝试解析数据 private func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] { guard buffer.readableBytes >= 2 else {