// // 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 writeFailed(Error) case internalError(Error) case timeout case decodeError(String) case packetTooLarge case dataStreamClosed } final class SDLQUICClient { enum State { case idle case running case stopped } private var state: State = .idle private let frameParser: SDLQUICFrameParser // 数据流 public var messageStream: AsyncThrowingStream private let messageCont: AsyncThrowingStream.Continuation private var isMessageContinuationFinished: Bool = false private var readTask: Task? private var connection: NWConnection? private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 专用队列保证线程安全 private let queueKey = DispatchSpecificKey() private let host: String private let port: UInt16 init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) { self.host = host self.port = port self.frameParser = SDLQUICFrameParser(maxBufferSize: maxBufferSize) (self.messageStream, self.messageCont) = AsyncThrowingStream.makeStream(of: SDLQUICInboundMessage.self) self.queue.setSpecific(key: self.queueKey, value: ()) } func start() { let options = NWProtocolTLS.Options() sec_protocol_options_add_tls_application_protocol( options.securityProtocolOptions, "punchnet/1.0" ) // 这里设置证书的校验逻辑 sec_protocol_options_set_verify_block( options.securityProtocolOptions, { metadata, trust, complete in // 执行公钥校验 complete(TLSVerifier.verify(trust: trust, host: self.host)) }, self.queue ) SDLLogger.log("[SDLQUICClient] start with tls protocol", for: .debug) let params = NWParameters(tls: options) // 关键:让 Network.framework 忽略系统代理 params.preferNoProxies = true let connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) connection.stateUpdateHandler = { [weak self] state in SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) switch state { case .ready: self?.startReadTask() self?.state = .running case .failed(let error): self?.finishMessageContinuationIfNeed(throwing: .connectionFailed(error)) case .cancelled: self?.finishMessageContinuationIfNeed(throwing: .connectionCancelled) default: () } } connection.start(queue: self.queue) self.connection = connection } private func finishMessageContinuationIfNeed(throwing error: SDLQUICError?) { if DispatchQueue.getSpecific(key: queueKey) != nil { self.finishMessageContinuationIfNeedOnQueue(throwing: error) } else { queue.async { [weak self] in self?.finishMessageContinuationIfNeedOnQueue(throwing: error) } } } private func finishMessageContinuationIfNeedOnQueue(throwing error: SDLQUICError?) { guard !self.isMessageContinuationFinished else { return } self.isMessageContinuationFinished = true if let error { self.messageCont.finish(throwing: error) } else { self.messageCont.finish() } } private func startReadTask() { self.readTask?.cancel() self.readTask = Task { do { while true { try Task.checkCancellation() let data = try await self.readOnce() let frames = try self.frameParser.parseFrames(data: data) for frame in frames { if let message = SDLQUICCodec.decode(frame: frame) { self.messageCont.yield(message) } else { self.finishMessageContinuationIfNeed(throwing: .decodeError("invalid message")) } } } } catch let err { self.finishMessageContinuationIfNeed(throwing: .internalError(err)) } } } func send(type: SDLPacketType, data: Data) { guard case .running = state, let connection = self.connection, connection.state == .ready else { return } 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) self?.finishMessageContinuationIfNeed(throwing: .writeFailed(error)) } }) } private func readOnce() async throws -> Data { guard let connection = self.connection else { throw SDLQUICError.connectionCancelled } return try await withCheckedThrowingContinuation { cont in connection.receive(minimumIncompleteLength: 1, maximumLength: 64 * 1024) { data, _, isComplete, error in if let error { cont.resume(throwing: error) return } if isComplete { cont.resume(throwing: SDLQUICError.dataStreamClosed) } else { cont.resume(returning: data ?? Data()) } } } } func stop() { if DispatchQueue.getSpecific(key: queueKey) != nil { self.stopOnQueue() } else { queue.sync { self.stopOnQueue() } } } private func stopOnQueue() { guard self.state != .stopped else { return } self.state = .stopped self.readTask?.cancel() self.connection?.cancel() self.finishMessageContinuationIfNeed(throwing: nil) } } // --MARK: 数据累加器 extension SDLQUICClient { class SDLQUICFrameParser { private let allocator = ByteBufferAllocator() // 最大缓冲区区为2M private let maxPacketSize: Int = 64 * 1024 private let maxBufferSize: Int private var buffer: ByteBuffer init(maxBufferSize: Int) { self.buffer = allocator.buffer(capacity: maxBufferSize) self.maxBufferSize = maxBufferSize } // 尝试解析数据 public func parseFrames(data: Data) throws -> [ByteBuffer] { self.buffer.writeBytes(data) 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) } } let threshold = maxBufferSize / 10 * 6 if buffer.readerIndex > threshold { buffer.discardReadBytes() } return frames } } } // --MARK: 编解码器 extension SDLQUICClient { enum SDLQUICCodec { public static 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: tls验证 extension SDLQUICClient { enum TLSVerifier { // 你的 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 } } } }