punchnet-macos/Tun/DNS/DNSLocalClient.swift
2026-05-28 14:34:11 +08:00

327 lines
11 KiB
Swift
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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)
}
private let queue = DispatchQueue(label: "com.sdl.DNSCloudClient.queue")
private let connection: NWConnection
private let timeoutInterval: TimeInterval = 3.0
nonisolated let packetFlow: AsyncThrowingStream<Data, Error>
private let packetContinuation: AsyncThrowingStream<Data, Error>.Continuation
private var isPacketContinuationFinished: Bool = false
private var pendingRequests: [UInt16: PendingRequest] = [:]
private var nextTransactionID: UInt16 = 1
private let readySignal = AsyncOneShot<Void>()
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)")
}
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()
let stream = Self.makeReceiveStream(for: self.connection)
for await data in stream {
try Task.checkCancellation()
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) {
guard connection.state == .ready, dnsPayload.count >= 2 else {
return
}
guard let allocatedTransactionID = self.allocateTransactionID() else {
SDLLogger.log("[DNSLocalClient] no available transaction id", for: .debug)
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")
}
private func handleConnectionStateUpdate(_ state: NWConnection.State) async {
switch state {
case .ready:
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 static func makeReceiveStream(for connection: NWConnection) -> AsyncStream<Data> {
return AsyncStream(bufferingPolicy: .bufferingNewest(256)) { continuation in
func receiveNext() {
connection.receiveMessage { content, _, _, error in
if let data = content, !data.isEmpty {
continuation.yield(data)
}
if error == nil && connection.state == .ready {
receiveNext()
} else {
continuation.finish()
}
}
}
receiveNext()
}
}
deinit {
SDLLogger.log("[DNSLocalClient] deinit", for: .debug)
}
}
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..<wordCount {
sum += UInt32(UInt16(bigEndian: words[i]))
}
if count % 2 != 0 {
let lastByte = ptr[count - 1]
sum += UInt32(lastByte) << 8
}
}
while (sum >> 16) != 0 {
sum = (sum & 0xffff) + (sum >> 16)
}
return UInt16(~sum & 0xffff)
}
}