359 lines
12 KiB
Swift
359 lines
12 KiB
Swift
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 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 run() async throws {
|
||
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)
|
||
}
|
||
}
|
||
|
||
self.connection = connection
|
||
|
||
connection.start(queue: .global())
|
||
|
||
try await withTaskCancellationHandler {
|
||
try await withThrowingTaskGroup(of: Void.self) { group in
|
||
defer {
|
||
group.cancelAll()
|
||
}
|
||
|
||
group.addTask { [weak self] in
|
||
let stream = Self.makeReceiveStream(for: connection)
|
||
for await data in stream {
|
||
try Task.checkCancellation()
|
||
guard let self else {
|
||
return
|
||
}
|
||
await self.handleResponse(data: data)
|
||
}
|
||
}
|
||
|
||
group.addTask { [weak self] in
|
||
while !Task.isCancelled {
|
||
try await Task.sleep(for: .seconds(3))
|
||
await self?.performCleanup()
|
||
}
|
||
}
|
||
|
||
try await group.next()
|
||
}
|
||
} onCancel: {
|
||
connection.cancel()
|
||
}
|
||
}
|
||
|
||
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 connection = self.connection
|
||
self.connection = nil
|
||
|
||
self.pendingRequests.removeAll()
|
||
self.nextTransactionID = 1
|
||
|
||
connection?.cancel()
|
||
self.finishPacketContinuationIfNeed(throwing: nil)
|
||
|
||
SDLLogger.log("[SDLLocalClient] stopped")
|
||
}
|
||
|
||
private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) {
|
||
switch state {
|
||
case .ready:
|
||
self.markConnectionReady(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 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) {
|
||
guard self.state != .stopped, self.isCurrentConnection(conn) else {
|
||
return
|
||
}
|
||
|
||
self.state = .running
|
||
}
|
||
|
||
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)
|
||
}
|
||
}
|