From 7b549981af572a5f6c2893f6a44a4ab9f06ee6b0 Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Tue, 28 Apr 2026 16:58:34 +0800 Subject: [PATCH] fix quicClient --- Tun/Punchnet/Actors/SDLQuicClient.swift | 291 ++++++++++++------------ 1 file changed, 146 insertions(+), 145 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLQuicClient.swift b/Tun/Punchnet/Actors/SDLQuicClient.swift index 58e5c17..0d44d2d 100644 --- a/Tun/Punchnet/Actors/SDLQuicClient.swift +++ b/Tun/Punchnet/Actors/SDLQuicClient.swift @@ -18,6 +18,7 @@ enum SDLQUICError: Error { case timeout case decodeError(String) case packetTooLarge + case waitReadyAlreadyInProgress } enum SDLQUICEvent: Error { @@ -32,7 +33,7 @@ enum SDLQUICClientExit: Error, Sendable, CustomStringConvertible { case transportClosed(String) case readFailed(String) case writeFailed(String) - + var description: String { switch self { case .normal: @@ -55,25 +56,35 @@ actor SDLQUICClient { private let maxPacketSize: Int // 最大缓冲区区为2M private let maxBufferSize: Int - - private let readyState = SDLQUICReadyState() - + + private enum ReadyStatus { + case idle + case connecting + case ready + case failed(Error) + case cancelled + } + + private var readyStatus: ReadyStatus = .idle + private var readyContinuation: CheckedContinuation? + private var readyTimeoutTask: Task? + // 消息流 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) { let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"]) - + self.maxBufferSize = maxBufferSize self.maxPacketSize = maxPacketSize (self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self) @@ -88,7 +99,7 @@ actor SDLQUICClient { }, self.queue ) - + let params = NWParameters(quic: options) self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) } @@ -101,27 +112,27 @@ actor SDLQUICClient { } connection.start(queue: self.queue) } - + private func handleConnectionStateUpdate(_ state: NWConnection.State) async { SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) switch state { case .ready: - await self.readyState.markReady() + self.markReady() case .failed(let error): - await self.readyState.markFailed(error) + self.markFailed(error) self.emitEvent(.failed(error)) case .cancelled: - await self.readyState.markCancelled() + self.markCancelled() self.emitEvent(.cancelled) case .setup, .preparing: - await self.readyState.markConnecting() + self.markConnecting() default: () } } - + func waitReady(timeout: Duration = .seconds(5)) async throws { - try await self.readyState.waitReady(timeout: timeout) + try await self.waitReadyUntilStateChanged(timeout: timeout) } func run() async -> SDLQUICClientExit { @@ -130,16 +141,16 @@ actor SDLQUICClient { group.addTask { await self.readLoop() } - + group.addTask { await self.heartbeatLoop() } - + let exit = await group.next() ?? .normal group.cancelAll() await self.stop() self.finishStreams() - + return exit } } onCancel: { @@ -148,14 +159,14 @@ actor SDLQUICClient { } } } - + func send(type: SDLPacketType, data: Data) { var len = UInt16(data.count + 1).bigEndian - + var packet = Data(Data(bytes: &len, count: 2)) packet.append(type.rawValue) packet.append(data) - + connection.send(content: packet, completion: .contentProcessed { [weak self] error in if let error { SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug) @@ -165,36 +176,36 @@ actor SDLQUICClient { } }) } - + 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 { self.connection.cancel() - await self.readyState.markCancelled() + self.markCancelled() self.finishStreams() } - + private func emitEvent(_ event: SDLQUICEvent) { guard !self.didFinishStreams else { return } - + self.eventCont.yield(event) } - + private func finishStreams() { guard !self.didFinishStreams else { return @@ -204,115 +215,105 @@ actor SDLQUICClient { self.messageCont.finish() self.eventCont.finish() } - + } // --MARK: Ready状态机 extension SDLQUICClient { - - actor SDLQUICReadyState { - enum State { - case idle - case connecting - case ready - case failed(Error) - case cancelled - } - private var state: State = .idle - private var continuations: [UUID: CheckedContinuation] = [:] + private func waitReadyUntilStateChanged(timeout: Duration) async throws { + try Task.checkCancellation() - func waitReady(timeout: Duration) async throws { - let id = UUID() - let timeoutTask = Task { - try? await Task.sleep(for: timeout) - if Task.isCancelled { - return - } - self.cancelWaiter(id: id, throwing: SDLQUICError.timeout) - } - - defer { - timeoutTask.cancel() - } - - try await withTaskCancellationHandler { - try await withCheckedThrowingContinuation { continuation in - self.addWaiter(id: id, continuation: continuation) - } - } onCancel: { - timeoutTask.cancel() - Task { - await self.cancelWaiter(id: id, throwing: CancellationError()) + try await withTaskCancellationHandler { + try await withCheckedThrowingContinuation { continuation in + switch self.readyStatus { + case .ready: + continuation.resume() + case .failed(let error): + continuation.resume(throwing: error) + case .cancelled: + continuation.resume(throwing: SDLQUICError.connectionCancelled) + case .idle, .connecting: + guard self.readyContinuation == nil else { + continuation.resume(throwing: SDLQUICError.waitReadyAlreadyInProgress) + return + } + self.readyContinuation = continuation + self.readyTimeoutTask?.cancel() + self.readyTimeoutTask = Task { + do { + try await Task.sleep(for: timeout) + if !Task.isCancelled { + await self.cancelReadyWaiter(throwing: SDLQUICError.timeout) + } + } catch { + return + } + } } } - } - - private func addWaiter(id: UUID, continuation: CheckedContinuation) { - switch state { - case .ready: - continuation.resume() - - case .failed(let error): - continuation.resume(throwing: error) - - case .cancelled: - continuation.resume(throwing: CancellationError()) - - case .idle, .connecting: - continuations[id] = continuation + } onCancel: { + Task { + await self.cancelReadyWaiter(throwing: CancellationError()) } } + } - private func cancelWaiter(id: UUID, throwing error: Error) { - guard let continuation = continuations.removeValue(forKey: id) else { - return - } - - continuation.resume(throwing: error) + private func cancelReadyWaiter(throwing error: Error) { + guard let continuation = self.readyContinuation else { + return } - func markConnecting() { - switch state { - case .idle: - state = .connecting - default: - break - } + self.readyContinuation = nil + self.readyTimeoutTask?.cancel() + self.readyTimeoutTask = nil + continuation.resume(throwing: error) + } + + private func markConnecting() { + switch self.readyStatus { + case .idle: + self.readyStatus = .connecting + default: + break + } + } + + private func markReady() { + self.readyStatus = .ready + self.resumeReadyWaiter() + } + + private func markFailed(_ error: Error) { + self.readyStatus = .failed(error) + self.resumeReadyWaiter(throwing: error) + } + + private func markCancelled() { + self.readyStatus = .cancelled + self.resumeReadyWaiter(throwing: SDLQUICError.connectionCancelled) + } + + private func resumeReadyWaiter() { + guard let continuation = self.readyContinuation else { + return } - func markReady() { - state = .ready + self.readyContinuation = nil + self.readyTimeoutTask?.cancel() + self.readyTimeoutTask = nil + continuation.resume() + } - let list = continuations - continuations.removeAll() - - for continuation in list.values { - continuation.resume() - } + private func resumeReadyWaiter(throwing error: Error) { + guard let continuation = self.readyContinuation else { + return } - func markFailed(_ error: Error) { - state = .failed(error) - - let list = continuations - continuations.removeAll() - - for continuation in list.values { - continuation.resume(throwing: error) - } - } - - func markCancelled() { - state = .cancelled - - let list = continuations - continuations.removeAll() - - for continuation in list.values { - continuation.resume(throwing: CancellationError()) - } - } + self.readyContinuation = nil + self.readyTimeoutTask?.cancel() + self.readyTimeoutTask = nil + continuation.resume(throwing: error) } } @@ -321,11 +322,11 @@ 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() @@ -335,14 +336,14 @@ extension SDLQUICClient { 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") } @@ -354,36 +355,36 @@ extension SDLQUICClient { return .readFailed("\(error)") } } - + // 尝试解析数据 private 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 { 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 @@ -406,7 +407,7 @@ extension SDLQUICClient { let packetType = SDLPacketType(rawValue: type) else { return nil } - + switch packetType { case .welcome: guard let bytes = buffer.readBytes(length: buffer.readableBytes), @@ -414,7 +415,7 @@ extension SDLQUICClient { return nil } return .welcome(welcome) - + case .registerSuperAck: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else { @@ -456,7 +457,7 @@ extension SDLQUICClient { return .pong default: SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug) - + return nil } } @@ -464,50 +465,50 @@ extension SDLQUICClient { // --MARK: quic验证 extension SDLQUICClient { - + enum QUICVerifier { // 你的 Base64 公钥指纹 static let pinnedPublicKeyHashes = [ "Q41r6hbMWEVyxo6heNAH4Wx/TH5NNOWlNif9bewcJ3E=" ] - + static func verify(trust: sec_trust_t, host: String) -> Bool { let secTrust = sec_trust_copy_ref(trust).takeRetainedValue() - + // --- Step 1: 系统验证 --- var error: CFError? guard SecTrustEvaluateWithError(secTrust, &error) else { SDLLogger.log("❌ 系统证书验证失败: \(error?.localizedDescription ?? "未知错误")", for: .debug) return false } - + // --- Step 2: 主机名验证 --- let policy = SecPolicyCreateSSL(true, host as CFString) SecTrustSetPolicies(secTrust, policy) - + guard SecTrustEvaluateWithError(secTrust, &error) else { SDLLogger.log("❌ 主机名校验失败: \(error?.localizedDescription ?? "未知错误")", for: .debug) return false } - + // --- Step 3: 获取叶子证书 --- guard let chain = SecTrustCopyCertificateChain(secTrust) as? [SecCertificate], let leafCertificate = chain.first else { SDLLogger.log("❌ 无法获取证书链或叶子证书", for: .debug) return false } - + // --- Step 4: 提取公钥 --- guard let publicKey = SecCertificateCopyKey(leafCertificate), let publicKeyData = SecKeyCopyExternalRepresentation(publicKey, nil) as Data? else { SDLLogger.log("❌ 无法提取公钥", for: .debug) return false } - + // --- Step 5: SHA256 校验 --- let hash = SHA256.hash(data: publicKeyData) let hashBase64 = Data(hash).base64EncodedString() - + if pinnedPublicKeyHashes.contains(hashBase64) { SDLLogger.log("✅ 公钥校验通过", for: .debug) return true @@ -516,6 +517,6 @@ extension SDLQUICClient { return false } } - + } }