punchnet-macos/Tun/Punchnet/DNS/DNSLocalClient.swift
2026-05-21 13:49:10 +08:00

374 lines
12 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
}
private enum State {
case idle
case starting
case running
case stopped
}
enum DNSLocalError: Error {
case failed(Error)
case cancelled
case sendFailed(Error)
}
private var state: State = .idle
private let dnsServerEndpoint: NWEndpoint
private var connection: NWConnection?
private var receiveTask: Task<Void, Never>?
private var cleanupTask: Task<Void, Never>?
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
init(host: String) {
self.dnsServerEndpoint = .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)")
}
}
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 start() {
guard self.state == .idle else {
return
}
self.state = .starting
let parameters = NWParameters.udp
parameters.prohibitedInterfaceTypes = [.other]
// 2. pathSelectionOptions
parameters.multipathServiceType = .handover
let connection = NWConnection(to: self.dnsServerEndpoint, using: parameters)
connection.stateUpdateHandler = { [weak self] state in
Task {
await self?.handleConnectionStateUpdate(state, for: connection)
}
}
let cleanupTask = Task { [weak self] in
while !Task.isCancelled {
try? await Task.sleep(nanoseconds: 3 * 1_000_000_000)
await self?.performCleanup()
}
}
self.connection = connection
self.cleanupTask = cleanupTask
connection.start(queue: .global())
}
func query(tracker: DNSTracker, dnsPayload: Data) {
guard self.state != .stopped,
let connection = self.connection, 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.state != .stopped else {
return
}
self.state = .stopped
let receiveTask = self.receiveTask
self.receiveTask = nil
let connection = self.connection
self.connection = nil
let cleanupTask = self.cleanupTask
self.cleanupTask = nil
self.pendingRequests.removeAll()
self.nextTransactionID = 1
receiveTask?.cancel()
connection?.cancel()
cleanupTask?.cancel()
self.finishPacketContinuationIfNeed(throwing: nil)
SDLLogger.log("[SDLLocalClient] stopped")
}
private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) {
switch state {
case .ready:
if self.markConnectionReady(conn) {
self.startReceiveTask(for: conn)
}
case .failed(let error):
SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug)
self.finishPacketContinuationIfNeed(throwing: .failed(error))
case .cancelled:
self.finishPacketContinuationIfNeed(throwing: .cancelled)
default:
()
}
}
private func startReceiveTask(for conn: NWConnection) {
let stream = Self.makeReceiveStream(for: conn)
let task = Task { [weak self] in
for await data in stream {
guard let self else {
break
}
await self.handleResponse(data: data)
}
}
let shouldKeepTask = self.state != .stopped && self.isCurrentConnection(conn)
if shouldKeepTask {
self.receiveTask?.cancel()
self.receiveTask = task
}
if !shouldKeepTask {
task.cancel()
}
}
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 markConnectionReady(_ conn: NWConnection) -> Bool {
guard self.state != .stopped, self.isCurrentConnection(conn) else {
return false
}
self.state = .running
return true
}
private func isCurrentConnection(_ conn: NWConnection) -> Bool {
guard let currentConnection = self.connection else {
return false
}
return currentConnection === conn
}
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 conn: NWConnection) -> AsyncStream<Data> {
return AsyncStream(bufferingPolicy: .bufferingNewest(256)) { continuation in
func receiveNext() {
conn.receiveMessage { content, _, _, error in
if let data = content, !data.isEmpty {
continuation.yield(data)
}
if error == nil && conn.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)
}
}