import Foundation import Network actor DNSLocalClient { struct DNSTracker { let transactionID: UInt16 let clientIP: UInt32 let clientPort: UInt16 let createdAt: Date } private struct PendingRequest { let tracker: DNSTracker } enum DNSLocalError: Error { case failed(Error) case cancelled case sendFailed(Error) case invalidData } private let queue = DispatchQueue(label: "com.sdl.DNSCloudClient.queue") private let connection: NWConnection private let timeoutInterval: TimeInterval = 3.0 nonisolated let packetFlow: AsyncThrowingStream private let packetContinuation: AsyncThrowingStream.Continuation private var isPacketContinuationFinished: Bool = false private var pendingRequests: [UInt16: PendingRequest] = [:] private var nextTransactionID: UInt16 = 1 private let readySignal = AsyncOneShot() private var isStopped: Bool = false init(host: String) { let dnsServerEndpoint = NWEndpoint.hostPort(host: Self.makeEndpointHost(ip: host), port: 53) let (stream, continuation) = AsyncThrowingStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256)) self.packetFlow = stream self.packetContinuation = continuation self.packetContinuation.onTermination = { termination in SDLLogger.log("[DNSLocalClient] packetFlow terminated: \(termination)", category: .dns) } let parameters = NWParameters.udp parameters.prohibitedInterfaceTypes = [.other] // 2. 增强健壮性:启用多路径切换(替代 pathSelectionOptions 的意图) parameters.multipathServiceType = .handover self.connection = NWConnection(to: dnsServerEndpoint, using: parameters) } private static func makeEndpointHost(ip: String) -> NWEndpoint.Host { if let ipv4Address = IPv4Address(ip) { return .ipv4(ipv4Address) } if let ipv6Address = IPv6Address(ip) { return .ipv6(ipv6Address) } preconditionFailure("invalid public DNS server IP: \(ip)") } func run() async throws { self.connection.stateUpdateHandler = { [weak self] state in Task { await self?.handleConnectionStateUpdate(state) } } self.connection.start(queue: self.queue) try await withTaskCancellationHandler { try await withThrowingTaskGroup(of: Void.self) { group in defer { group.cancelAll() } group.addTask { try await self.readySignal.wait() while true { try Task.checkCancellation() let data = try await self.readOnce() await self.handleResponse(data: data) } } group.addTask { while !Task.isCancelled { try await Task.sleep(for: .seconds(3)) await self.performCleanup() } } try await group.next() } } onCancel: { self.connection.cancel() } } func query(tracker: DNSTracker, dnsPayload: Data) async { guard dnsPayload.count >= 2 else { return } do { try await self.readySignal.wait(timeout: .seconds(3)) } catch { SDLLogger.log("[DNSLocalClient] drop query before ready: \(error)", category: .dns) return } guard !self.isStopped, connection.state == .ready else { return } guard let allocatedTransactionID = self.allocateTransactionID() else { SDLLogger.log("[DNSLocalClient] no available transaction id", category: .dns) return } let transactionID = allocatedTransactionID self.pendingRequests[transactionID] = PendingRequest(tracker: tracker) let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID) connection.send(content: rewrittenPayload, completion: .contentProcessed { [weak self] error in if let error { Task { await self?.handleSendFailure(transactionID: transactionID, error: error) } } }) } func stop() { guard !self.isStopped else { return } self.isStopped = true self.connection.cancel() self.pendingRequests.removeAll() self.nextTransactionID = 1 self.finishPacketContinuationIfNeed(throwing: nil) SDLLogger.log("[SDLLocalClient] stopped", category: .dns) } private func handleConnectionStateUpdate(_ state: NWConnection.State) async { switch state { case .ready: guard !self.isStopped else { return } await self.readySignal.succeed(()) case .failed(let error): await self.readySignal.fail(DNSLocalError.failed(error)) self.finishPacketContinuationIfNeed(throwing: .failed(error)) case .cancelled: await self.readySignal.fail(DNSLocalError.cancelled) self.finishPacketContinuationIfNeed(throwing: .cancelled) default: () } } private func finishPacketContinuationIfNeed(throwing error: DNSLocalError?) { guard !self.isPacketContinuationFinished else { return } self.isPacketContinuationFinished = true if let error { self.packetContinuation.finish(throwing: error) } else { self.packetContinuation.finish() } } private func handleResponse(data: Data) { guard let rewrittenTransactionID = Self.readTransactionID(from: data), let pendingRequest = self.pendingRequests.removeValue(forKey: rewrittenTransactionID) else { return } let restoredPayload = Self.rewriteTransactionID(in: data, to: pendingRequest.tracker.transactionID) let packet = Self.createDNSResponse( payload: restoredPayload, srcIP: DNSHelper.dnsDestIpAddr, srcPort: 53, destIP: pendingRequest.tracker.clientIP, destPort: pendingRequest.tracker.clientPort ) self.packetContinuation.yield(packet) } private func handleSendFailure(transactionID: UInt16, error: NWError) { self.pendingRequests.removeValue(forKey: transactionID) self.finishPacketContinuationIfNeed(throwing: .sendFailed(error)) } private func performCleanup() { let now = Date() self.pendingRequests = self.pendingRequests.filter { _, request in now.timeIntervalSince(request.tracker.createdAt) < self.timeoutInterval } } private func allocateTransactionID() -> UInt16? { var candidate = self.nextTransactionID == 0 ? 1 : self.nextTransactionID let start = candidate repeat { if self.pendingRequests[candidate] == nil { self.nextTransactionID = Self.nextTransactionID(after: candidate) return candidate } candidate = Self.nextTransactionID(after: candidate) } while candidate != start return nil } private static func nextTransactionID(after id: UInt16) -> UInt16 { return id == UInt16.max ? 1 : id &+ 1 } private static func readTransactionID(from payload: Data) -> UInt16? { guard payload.count >= 2 else { return nil } return UInt16(payload[0]) << 8 | UInt16(payload[1]) } private static func rewriteTransactionID(in payload: Data, to transactionID: UInt16) -> Data { guard payload.count >= 2 else { return payload } var rewrittenPayload = payload rewrittenPayload[0] = UInt8((transactionID >> 8) & 0xFF) rewrittenPayload[1] = UInt8(transactionID & 0xFF) return rewrittenPayload } private func readOnce() async throws -> Data { guard self.connection.state == .ready else { throw DNSLocalError.cancelled } let readContinuation = OnceContinuation() return try await withTaskCancellationHandler { try await withCheckedThrowingContinuation { cont in readContinuation.set(cont) self.connection.receiveMessage { content, _, _, error in if let error { readContinuation.resume(throwing: error) } else if let data = content, !data.isEmpty { readContinuation.resume(returning: data) } else { readContinuation.resume(throwing: DNSLocalError.invalidData) } } } } onCancel: { readContinuation.resume(throwing: CancellationError()) } } deinit { SDLLogger.log("[DNSLocalClient] deinit", category: .dns) } } extension DNSLocalClient { static func createDNSResponse(payload: Data, srcIP: UInt32, srcPort: UInt16, destIP: UInt32, destPort: UInt16) -> Data { let udpLen = 8 + payload.count let ipLen = 20 + udpLen var ipHeader = Data(count: 20) ipHeader[0] = 0x45 ipHeader[2...3] = withUnsafeBytes(of: UInt16(ipLen).bigEndian) { Data($0) } ipHeader[8] = 64 ipHeader[9] = 17 ipHeader[12...15] = withUnsafeBytes(of: srcIP.bigEndian) { Data($0) } ipHeader[16...19] = withUnsafeBytes(of: destIP.bigEndian) { Data($0) } let ipChecksum = calculateChecksum(data: ipHeader) ipHeader[10...11] = withUnsafeBytes(of: ipChecksum.bigEndian) { Data($0) } var udpHeader = Data(count: 8) udpHeader[0...1] = withUnsafeBytes(of: srcPort.bigEndian) { Data($0) } udpHeader[2...3] = withUnsafeBytes(of: destPort.bigEndian) { Data($0) } udpHeader[4...5] = withUnsafeBytes(of: UInt16(udpLen).bigEndian) { Data($0) } udpHeader[6...7] = Data([0, 0]) var packet = Data(capacity: ipLen) packet.append(ipHeader) packet.append(udpHeader) packet.append(payload) return packet } static func calculateChecksum(data: Data) -> UInt16 { var sum: UInt32 = 0 let count = data.count data.withUnsafeBytes { (ptr: UnsafeRawBufferPointer) in guard let baseAddress = ptr.baseAddress else { return } let wordCount = count / 2 let words = baseAddress.bindMemory(to: UInt16.self, capacity: wordCount) for i in 0..> 16) != 0 { sum = (sum & 0xffff) + (sum >> 16) } return UInt16(~sum & 0xffff) } }