// // SDLQuicClient.swift // Tun // // Created by 安礼成 on 2026/2/13. // import Foundation import NIOCore import Network import CryptoKit import Security // 定义错误类型,便于上层处理 enum SDLQUICError: Error { case connectionFailed(Error) case connectionCancelled case timeout case decodeError(String) case packetTooLarge } enum SDLQUICEvent: Error { case failed(Error) case cancelled 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 { private let allocator = ByteBufferAllocator() // 单个包最大64K private let maxPacketSize: Int // 最大缓冲区区为2M private let maxBufferSize: Int private let readyState = SDLQUICReadyState() // 消息流 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) (self.eventStream, self.eventCont) = AsyncStream.makeStream(of: SDLQUICEvent.self) // TODO 这里设置证书的校验逻辑 sec_protocol_options_set_verify_block( options.securityProtocolOptions, { metadata, trust, complete in // 执行公钥校验 complete(QUICVerifier.verify(trust: trust, host: host)) }, self.queue ) let params = NWParameters(quic: options) self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) } func start() { connection.stateUpdateHandler = { [weak self] state in Task { await self?.handleConnectionStateUpdate(state) } } 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() case .failed(let error): await self.readyState.markFailed(error) self.emitEvent(.failed(error)) case .cancelled: await self.readyState.markCancelled() self.emitEvent(.cancelled) case .setup, .preparing: await self.readyState.markConnecting() default: () } } func waitReady(timeout: Duration = .seconds(5)) async throws { try await self.readyState.waitReady(timeout: timeout) } func run() async -> SDLQUICClientExit { await withTaskCancellationHandler { 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() await self.stop() self.finishStreams() return exit } } onCancel: { Task { await self.stop() } } } 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) Task { await self?.emitEvent(.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 { self.connection.cancel() await self.readyState.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 } self.didFinishStreams = true self.messageCont.finish() self.eventCont.finish() } deinit { self.connection.cancel() self.finishStreams() } } // --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] = [:] func waitReady(timeout: Duration) async throws { let id = UUID() let timeoutTask = Task { try? await Task.sleep(for: timeout) if Task.isCancelled { return } await 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()) } } } 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 } } private func cancelWaiter(id: UUID, throwing error: Error) { guard let continuation = continuations.removeValue(forKey: id) else { return } continuation.resume(throwing: error) } func markConnecting() { switch state { case .idle: state = .connecting default: break } } func markReady() { state = .ready let list = continuations continuations.removeAll() for continuation in list.values { continuation.resume() } } 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()) } } } } // --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() } 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 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 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? { var buffer = frame guard let type = buffer.readInteger(as: UInt8.self), let packetType = SDLPacketType(rawValue: type) else { return nil } switch packetType { case .welcome: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let welcome = try? SDLWelcome(serializedBytes: bytes) else { return nil } return .welcome(welcome) case .registerSuperAck: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else { return nil } return .registerSuperAck(registerSuperAck) case .registerSuperNak: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let registerSuperNak = try? SDLRegisterSuperNak(serializedBytes: bytes) else { return nil } return .registerSuperNak(registerSuperNak) case .peerInfo: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let peerInfo = try? SDLPeerInfo(serializedBytes: bytes) else { return nil } return .peerInfo(peerInfo) case .policyResponse: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let policyResponse = try? SDLPolicyResponse(serializedBytes: bytes) else { return nil } return .policyReponse(policyResponse) case .arpResponse: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let arpResponse = try? SDLArpResponse(serializedBytes: bytes) else { return nil } return .arpResponse(arpResponse) case .event: guard let bytes = buffer.readBytes(length: buffer.readableBytes), let event = try? SDLEvent(serializedBytes: bytes) else { SDLLogger.log("SDLQUICClient decode Event Error", for: .debug) return nil } return .event(event) case .pong: return .pong default: SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug) return nil } } } // --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 } else { SDLLogger.log("⚠️ 公钥不匹配! 收到: \(hashBase64)", for: .debug) return false } } } }