From dde40ceb68f1109abab4f8ca9f984ef00086b58b Mon Sep 17 00:00:00 2001 From: anlicheng <244108715@qq.com> Date: Tue, 5 May 2026 12:41:16 +0800 Subject: [PATCH] fix dnsClient --- Tun/Punchnet/Actors/SDLContextActor.swift | 56 +++------ Tun/Punchnet/DNS/DNSLocalClient.swift | 131 +++++++++++++++------- 2 files changed, 107 insertions(+), 80 deletions(-) diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index d86b217..e1dc13a 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -341,49 +341,23 @@ actor SDLContextActor { SDLLogger.log("[SDLContext] dnsLocalClient started") self.dnsLocalClient = dnsLocalClient - try await withThrowingTaskGroup { group in - defer { - group.cancelAll() + defer { + self.dnsLocalClient = nil + //dnsLocalClient.stop() + } + + try await withTaskCancellationHandler { + // 处理事件流 + for try await packet in dnsLocalClient.packetFlow { + try Task.checkCancellation() + // 要想办法构造一个完整的Ip包 + let nePacket = NEPacket(data: packet, protocolFamily: 2) + self.provider.packetFlow.writePacketObjects([nePacket]) } - - group.addTask { - // 处理事件流 - for await packet in dnsLocalClient.packetFlow { - try Task.checkCancellation() - - // 要想办法构造一个完整的Ip包 - let nePacket = NEPacket(data: packet, protocolFamily: 2) - self.provider.packetFlow.writePacketObjects([nePacket]) - } - throw SDLContextError.dnsLocalClientClosed - } - - 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") - throw SDLContextError.dnsLocalClientCancelled - case .sendFailed(let error): - SDLLogger.log("[SDLContext] dnsLocalClient sendFailed: \(error)") - throw error - } - } - } - - do { - try await group.next() + throw SDLContextError.dnsLocalClientClosed + } onCancel: { + Task { await dnsLocalClient.stop() - self.dnsLocalClient = nil - } catch let err { - await dnsLocalClient.stop() - self.dnsLocalClient = nil - throw err } } } diff --git a/Tun/Punchnet/DNS/DNSLocalClient.swift b/Tun/Punchnet/DNS/DNSLocalClient.swift index 3be58b4..f1b8984 100644 --- a/Tun/Punchnet/DNS/DNSLocalClient.swift +++ b/Tun/Punchnet/DNS/DNSLocalClient.swift @@ -16,11 +16,12 @@ actor DNSLocalClient { private enum State { case idle + case starting case running case stopped } - enum Event { + enum DNSLocalError: Error { case failed(Error) case cancelled case sendFailed(Error) @@ -35,33 +36,32 @@ actor DNSLocalClient { private var cleanupTask: Task? private let timeoutInterval: TimeInterval = 3.0 - - - public let packetFlow: AsyncStream + nonisolated let packetFlow: AsyncThrowingStream @ObservationIgnored - private let packetContinuation: AsyncStream.Continuation - - // 事件处理 - public let eventStream: AsyncStream - @ObservationIgnored - private let eventContinuation: AsyncStream.Continuation + private let packetContinuation: AsyncThrowingStream.Continuation + private var isPacketContinuationFinished: Bool = false private var pendingRequests: [UInt16: PendingRequest] = [:] private var nextTransactionID: UInt16 = 1 - + init(host: String) { self.dnsServerEndpoint = .hostPort(host: NWEndpoint.Host(host), port: 53) - let (stream, continuation) = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256)) + let (stream, continuation) = AsyncThrowingStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256)) self.packetFlow = stream self.packetContinuation = continuation - let eventPair = AsyncStream.makeStream(of: Event.self) - self.eventStream = eventPair.stream - self.eventContinuation = eventPair.continuation + self.packetContinuation.onTermination = { termination in + SDLLogger.log("[DNSLocalClient] packetFlow terminated: \(termination)") + } } func start() { + guard self.state == .idle else { + return + } + self.state = .starting + let parameters = NWParameters.udp parameters.prohibitedInterfaceTypes = [.other] // 2. 增强健壮性:启用多路径切换(替代 pathSelectionOptions 的意图) @@ -73,44 +73,47 @@ actor DNSLocalClient { await self?.handleConnectionStateUpdate(state, for: connection) } } - connection.start(queue: .global()) - self.connection = connection - - self.cleanupTask = Task { [weak self] in + + 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 let connection = self.connection, connection.state == .ready, dnsPayload.count >= 2 else { + let transactionID: UInt16 + let connection: NWConnection + let rewrittenPayload: Data + + guard self.state != .stopped, + let connection = self.connection, connection.state == .ready, dnsPayload.count >= 2 else { return } - guard let transactionID = self.allocateTransactionID() else { + guard let allocatedTransactionID = self.allocateTransactionID() else { SDLLogger.log("[DNSLocalClient] no available transaction id", for: .debug) return } + transactionID = allocatedTransactionID self.pendingRequests[transactionID] = PendingRequest(tracker: tracker) - let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID) + rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID) - connection.send(content: rewrittenPayload, completion: .contentProcessed { error in + connection.send(content: rewrittenPayload, completion: .contentProcessed { [weak self] error in if let error { - self.eventContinuation.yield(.sendFailed(error)) Task { - await self.removePendingRequest(forKey: transactionID) + await self?.handleSendFailure(transactionID: transactionID, error: error) } } }) } - private func removePendingRequest(forKey id: UInt16) { - self.pendingRequests.removeValue(forKey: id) - } - func stop() { guard self.state != .stopped else { return @@ -118,31 +121,35 @@ actor DNSLocalClient { self.state = .stopped - self.receiveTask?.cancel() + let receiveTask = self.receiveTask self.receiveTask = nil - self.connection?.cancel() + let connection = self.connection self.connection = nil - self.cleanupTask?.cancel() + let cleanupTask = self.cleanupTask self.cleanupTask = nil self.pendingRequests.removeAll() self.nextTransactionID = 1 - - self.packetContinuation.finish() + + receiveTask?.cancel() + connection?.cancel() + cleanupTask?.cancel() + self.finishPacketContinuationIfNeed(throwing: nil) } private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) { switch state { case .ready: - self.startReceiveTask(for: conn) - self.state = .running + if self.markConnectionReady(conn) { + self.startReceiveTask(for: conn) + } case .failed(let error): SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug) - self.eventContinuation.yield(.failed(error)) + self.finishPacketContinuationIfNeed(throwing: .failed(error)) case .cancelled: - self.eventContinuation.yield(.cancelled) + self.finishPacketContinuationIfNeed(throwing: .cancelled) default: () } @@ -151,7 +158,7 @@ actor DNSLocalClient { private func startReceiveTask(for conn: NWConnection) { let stream = Self.makeReceiveStream(for: conn) - self.receiveTask = Task { [weak self] in + let task = Task { [weak self] in for await data in stream { guard let self else { break @@ -159,6 +166,30 @@ actor DNSLocalClient { 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) { @@ -178,6 +209,28 @@ actor DNSLocalClient { ) 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()