fix dnsLocalClient

This commit is contained in:
anlicheng 2026-05-04 22:04:12 +08:00
parent 288c58d1a7
commit 2505945825
3 changed files with 115 additions and 112 deletions

View File

@ -58,8 +58,8 @@ actor SDLContextActor {
private var dnsClient: DNSCloudClient? private var dnsClient: DNSCloudClient?
// Localdnsclient // Localdnsclient
private let publicDnsServers = ["223.5.5.5", "119.29.29.29"]
private var dnsLocalClient: DNSLocalClient? private var dnsLocalClient: DNSLocalClient?
private var dnsLocalWorker: Task<Void, Never>?
private var quicClient: SDLQUICClient? private var quicClient: SDLQUICClient?
@ -351,28 +351,58 @@ actor SDLContextActor {
} }
private func startDnsLocalClient() async { private func startDnsLocalClient() async throws {
self.dnsLocalWorker?.cancel() let dnsServer = self.publicDnsServers.randomElement() ?? self.publicDnsServers[0]
self.dnsLocalWorker = nil
// dns // dns
let dnsLocalClient = DNSLocalClient() let dnsLocalClient = DNSLocalClient(host: dnsServer)
await dnsLocalClient.start() await dnsLocalClient.start()
SDLLogger.log("[SDLContext] dnsClient started") SDLLogger.log("[SDLContext] dnsLocalClient started")
self.dnsLocalClient = dnsLocalClient self.dnsLocalClient = dnsLocalClient
let packetFlow = dnsLocalClient.packetFlow
self.dnsLocalWorker = Task.detached { try await withThrowingTaskGroup { group in
// defer {
for await packet in packetFlow { group.cancelAll()
if Task.isCancelled {
break
} }
group.addTask {
//
for await packet in dnsLocalClient.packetFlow {
try Task.checkCancellation()
// Ip // Ip
let nePacket = NEPacket(data: packet, protocolFamily: 2) let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket]) self.provider.packetFlow.writePacketObjects([nePacket])
} }
} }
group.addTask {
for await event in dnsLocalClient.eventStream {
try Task.checkCancellation()
switch event {
case .failed(let error):
SDLLogger.log("[SDLContext] dnsLocalClient failed: \(error)")
throw error
case .cancelled:
SDLLogger.log("[SDLContext] dnsLocalClient cancelled")
return
case .sendFailed(let error):
SDLLogger.log("[SDLContext] dnsLocalClient sendFailed: \(error)")
throw error
}
}
}
do {
try await group.next()
await self.dnsLocalClient?.stop()
self.dnsLocalClient = nil
} catch let err {
await self.dnsLocalClient?.stop()
self.dnsLocalClient = nil
throw err
}
}
} }
private func startUDPHole() async throws { private func startUDPHole() async throws {
@ -498,11 +528,6 @@ actor SDLContextActor {
// context // context
public func stop() async { public func stop() async {
await self.stopRuntime()
}
private func stopRuntime() async {
await self.supervisor.stop() await self.supervisor.stop()
await self.puncherActor.stop() await self.puncherActor.stop()
await self.arpServer.clear() await self.arpServer.clear()
@ -520,14 +545,10 @@ actor SDLContextActor {
self.quicClient?.stop() self.quicClient?.stop()
self.quicClient = nil self.quicClient = nil
await self.dnsClient?.stop() self.dnsClient?.stop()
self.dnsWorker?.cancel()
self.dnsWorker = nil
self.dnsClient = nil self.dnsClient = nil
await self.dnsLocalClient?.stop() await self.dnsLocalClient?.stop()
self.dnsLocalWorker?.cancel()
self.dnsLocalWorker = nil
self.dnsLocalClient = nil self.dnsLocalClient = nil
self.readTask?.cancel() self.readTask?.cancel()

View File

@ -20,48 +20,61 @@ actor DNSLocalClient {
case stopped case stopped
} }
private var state: State = .idle enum Event {
private var connections: [NWConnection] = [] case failed(Error)
private var receiveTasks: [ObjectIdentifier: Task<Void, Never>] = [:] case cancelled
private let dnsServers = ["223.5.5.5", "119.29.29.29"] case sendFailed(Error)
}
let packetFlow: AsyncStream<Data> 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
public let packetFlow: AsyncStream<Data>
@ObservationIgnored
private let packetContinuation: AsyncStream<Data>.Continuation private let packetContinuation: AsyncStream<Data>.Continuation
//
public let eventStream: AsyncStream<Event>
@ObservationIgnored
private let eventContinuation: AsyncStream<Event>.Continuation
private var pendingRequests: [UInt16: PendingRequest] = [:] private var pendingRequests: [UInt16: PendingRequest] = [:]
private var nextTransactionID: UInt16 = 1 private var nextTransactionID: UInt16 = 1
private var cleanupTask: Task<Void, Never>? init(host: String) {
private let timeoutInterval: TimeInterval = 3.0 self.dnsServerEndpoint = .hostPort(host: NWEndpoint.Host(host), port: 53)
private var didFinishPacketFlow = false
init() {
let (stream, continuation) = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256)) let (stream, continuation) = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256))
self.packetFlow = stream self.packetFlow = stream
self.packetContinuation = continuation self.packetContinuation = continuation
let eventPair = AsyncStream.makeStream(of: Event.self)
self.eventStream = eventPair.stream
self.eventContinuation = eventPair.continuation
} }
func start() { func start() {
guard case .idle = self.state else {
return
}
self.state = .running
for server in self.dnsServers {
let endpoint = NWEndpoint.hostPort(host: NWEndpoint.Host(server), port: 53)
let parameters = NWParameters.udp let parameters = NWParameters.udp
parameters.prohibitedInterfaceTypes = [.other] parameters.prohibitedInterfaceTypes = [.other]
// 2. pathSelectionOptions
parameters.multipathServiceType = .handover
let conn = NWConnection(to: endpoint, using: parameters) let connection = NWConnection(to: self.dnsServerEndpoint, using: parameters)
conn.stateUpdateHandler = { [weak self] state in connection.stateUpdateHandler = { [weak self] state in
Task { Task {
await self?.handleConnectionStateUpdate(state, for: conn) await self?.handleConnectionStateUpdate(state, for: connection)
} }
} }
conn.start(queue: .global()) connection.start(queue: .global())
self.connections.append(conn) self.connection = connection
}
self.cleanupTask = Task { [weak self] in self.cleanupTask = Task { [weak self] in
while !Task.isCancelled { while !Task.isCancelled {
@ -72,7 +85,7 @@ actor DNSLocalClient {
} }
func query(tracker: DNSTracker, dnsPayload: Data) { func query(tracker: DNSTracker, dnsPayload: Data) {
guard case .running = self.state, dnsPayload.count >= 2 else { guard let connection = self.connection, connection.state == .ready, dnsPayload.count >= 2 else {
return return
} }
@ -84,19 +97,18 @@ actor DNSLocalClient {
self.pendingRequests[transactionID] = PendingRequest(tracker: tracker) self.pendingRequests[transactionID] = PendingRequest(tracker: tracker)
let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID) let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID)
var hasReadyConnection = false connection.send(content: rewrittenPayload, completion: .contentProcessed { error in
for conn in self.connections where conn.state == .ready {
hasReadyConnection = true
conn.send(content: rewrittenPayload, completion: .contentProcessed({ error in
if let error { if let error {
SDLLogger.log("[DNSLocalClient] send error: \(error.localizedDescription)", for: .debug) self.eventContinuation.yield(.sendFailed(error))
Task {
await self.removePendingRequest(forKey: transactionID)
} }
})) }
})
} }
if !hasReadyConnection { private func removePendingRequest(forKey id: UInt16) {
self.pendingRequests.removeValue(forKey: transactionID) self.pendingRequests.removeValue(forKey: id)
}
} }
func stop() { func stop() {
@ -105,69 +117,52 @@ actor DNSLocalClient {
} }
self.state = .stopped self.state = .stopped
self.receiveTasks.values.forEach { $0.cancel() }
self.receiveTasks.removeAll() self.receiveTask?.cancel()
self.connections.forEach { $0.cancel() } self.receiveTask = nil
self.connections.removeAll()
self.connection?.cancel()
self.connection = nil
self.cleanupTask?.cancel() self.cleanupTask?.cancel()
self.cleanupTask = nil self.cleanupTask = nil
self.pendingRequests.removeAll() self.pendingRequests.removeAll()
self.nextTransactionID = 1 self.nextTransactionID = 1
self.finishPacketFlowIfNeeded()
self.packetContinuation.finish()
} }
private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) { private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) {
guard case .running = self.state else {
return
}
switch state { switch state {
case .ready: case .ready:
self.startReceiveTask(for: conn) self.startReceiveTask(for: conn)
self.state = .running
case .failed(let error): case .failed(let error):
SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug) SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug)
self.stop() self.eventContinuation.yield(.failed(error))
case .cancelled: case .cancelled:
let key = ObjectIdentifier(conn) self.eventContinuation.yield(.cancelled)
self.receiveTasks.removeValue(forKey: key)?.cancel()
self.connections.removeAll { $0 === conn }
if self.connections.isEmpty {
self.stop()
}
default: default:
() ()
} }
} }
private func startReceiveTask(for conn: NWConnection) { private func startReceiveTask(for conn: NWConnection) {
let key = ObjectIdentifier(conn)
guard self.receiveTasks[key] == nil else {
return
}
let stream = Self.makeReceiveStream(for: conn) let stream = Self.makeReceiveStream(for: conn)
self.receiveTasks[key] = Task { [weak self] in
self.receiveTask = Task { [weak self] in
for await data in stream { for await data in stream {
guard let self else { guard let self else {
break break
} }
await self.handleResponse(data: data) await self.handleResponse(data: data)
} }
await self?.didFinishReceiving(for: conn)
} }
} }
private func didFinishReceiving(for conn: NWConnection) {
let key = ObjectIdentifier(conn)
self.receiveTasks.removeValue(forKey: key)
}
private func handleResponse(data: Data) { private func handleResponse(data: Data) {
guard case .running = self.state, guard let rewrittenTransactionID = Self.readTransactionID(from: data),
let rewrittenTransactionID = Self.readTransactionID(from: data),
let pendingRequest = self.pendingRequests.removeValue(forKey: rewrittenTransactionID) else { let pendingRequest = self.pendingRequests.removeValue(forKey: rewrittenTransactionID) else {
return return
} }
@ -185,10 +180,6 @@ actor DNSLocalClient {
} }
private func performCleanup() { private func performCleanup() {
guard case .running = self.state else {
return
}
let now = Date() let now = Date()
self.pendingRequests = self.pendingRequests.filter { _, request in self.pendingRequests = self.pendingRequests.filter { _, request in
now.timeIntervalSince(request.tracker.createdAt) < self.timeoutInterval now.timeIntervalSince(request.tracker.createdAt) < self.timeoutInterval
@ -211,15 +202,6 @@ actor DNSLocalClient {
return nil return nil
} }
private func finishPacketFlowIfNeeded() {
guard !self.didFinishPacketFlow else {
return
}
self.didFinishPacketFlow = true
self.packetContinuation.finish()
}
private static func nextTransactionID(after id: UInt16) -> UInt16 { private static func nextTransactionID(after id: UInt16) -> UInt16 {
return id == UInt16.max ? 1 : id &+ 1 return id == UInt16.max ? 1 : id &+ 1
} }