diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 901c397..8bbdad9 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -58,8 +58,8 @@ actor SDLContextActor { private var dnsClient: DNSCloudClient? // Localdns的client对象 + private let publicDnsServers = ["223.5.5.5", "119.29.29.29"] private var dnsLocalClient: DNSLocalClient? - private var dnsLocalWorker: Task? private var quicClient: SDLQUICClient? @@ -351,26 +351,56 @@ actor SDLContextActor { } - private func startDnsLocalClient() async { - self.dnsLocalWorker?.cancel() - self.dnsLocalWorker = nil - + private func startDnsLocalClient() async throws { + let dnsServer = self.publicDnsServers.randomElement() ?? self.publicDnsServers[0] // 启动dns服务 - let dnsLocalClient = DNSLocalClient() + let dnsLocalClient = DNSLocalClient(host: dnsServer) await dnsLocalClient.start() - SDLLogger.log("[SDLContext] dnsClient started") + SDLLogger.log("[SDLContext] dnsLocalClient started") self.dnsLocalClient = dnsLocalClient - let packetFlow = dnsLocalClient.packetFlow - self.dnsLocalWorker = Task.detached { - // 处理事件流 - for await packet in packetFlow { - if Task.isCancelled { - break + + try await withThrowingTaskGroup { group in + defer { + group.cancelAll() + } + + group.addTask { + // 处理事件流 + for await packet in dnsLocalClient.packetFlow { + try Task.checkCancellation() + + // 要想办法构造一个完整的Ip包 + let nePacket = NEPacket(data: packet, protocolFamily: 2) + self.provider.packetFlow.writePacketObjects([nePacket]) } - - // 要想办法构造一个完整的Ip包 - let nePacket = NEPacket(data: packet, protocolFamily: 2) - 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 } } } @@ -498,11 +528,6 @@ actor SDLContextActor { // 处理context的停止问题 public func stop() async { - await self.stopRuntime() - } - - private func stopRuntime() async { - await self.supervisor.stop() await self.puncherActor.stop() await self.arpServer.clear() @@ -520,14 +545,10 @@ actor SDLContextActor { self.quicClient?.stop() self.quicClient = nil - await self.dnsClient?.stop() - self.dnsWorker?.cancel() - self.dnsWorker = nil + self.dnsClient?.stop() self.dnsClient = nil await self.dnsLocalClient?.stop() - self.dnsLocalWorker?.cancel() - self.dnsLocalWorker = nil self.dnsLocalClient = nil self.readTask?.cancel() diff --git a/Tun/Punchnet/DNS/DNSCloudClient.swift b/Tun/Punchnet/DNS/DNSCloudClient.swift index 9ab242a..f068a30 100644 --- a/Tun/Punchnet/DNS/DNSCloudClient.swift +++ b/Tun/Punchnet/DNS/DNSCloudClient.swift @@ -37,7 +37,7 @@ final class DNSCloudClient { /// - Parameter host: 你的 sn-server 地址 (如 "8.8.8.8") /// - Parameter port: 端口 (如 53) - init(host: String, port: UInt16 ) { + init(host: String, port: UInt16) { self.dnsServerAddress = .hostPort(host: NWEndpoint.Host(host), port: NWEndpoint.Port(integerLiteral: port)) let packetPair = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256)) diff --git a/Tun/Punchnet/DNS/DNSLocalClient.swift b/Tun/Punchnet/DNS/DNSLocalClient.swift index 6aef9b6..3be58b4 100644 --- a/Tun/Punchnet/DNS/DNSLocalClient.swift +++ b/Tun/Punchnet/DNS/DNSLocalClient.swift @@ -20,48 +20,61 @@ actor DNSLocalClient { case stopped } + enum Event { + case failed(Error) + case cancelled + case sendFailed(Error) + } + private var state: State = .idle - private var connections: [NWConnection] = [] - private var receiveTasks: [ObjectIdentifier: Task] = [:] - private let dnsServers = ["223.5.5.5", "119.29.29.29"] - let packetFlow: AsyncStream - private let packetContinuation: AsyncStream.Continuation - - private var pendingRequests: [UInt16: PendingRequest] = [:] - private var nextTransactionID: UInt16 = 1 + private let dnsServerEndpoint: NWEndpoint + private var connection: NWConnection? + private var receiveTask: Task? private var cleanupTask: Task? private let timeoutInterval: TimeInterval = 3.0 - private var didFinishPacketFlow = false - init() { + + + public let packetFlow: AsyncStream + @ObservationIgnored + private let packetContinuation: AsyncStream.Continuation + + // 事件处理 + public let eventStream: AsyncStream + @ObservationIgnored + private let eventContinuation: AsyncStream.Continuation + + 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)) self.packetFlow = stream self.packetContinuation = continuation + + let eventPair = AsyncStream.makeStream(of: Event.self) + self.eventStream = eventPair.stream + self.eventContinuation = eventPair.continuation } func start() { - guard case .idle = self.state else { - return - } + let parameters = NWParameters.udp + parameters.prohibitedInterfaceTypes = [.other] + // 2. 增强健壮性:启用多路径切换(替代 pathSelectionOptions 的意图) + parameters.multipathServiceType = .handover - self.state = .running - - for server in self.dnsServers { - let endpoint = NWEndpoint.hostPort(host: NWEndpoint.Host(server), port: 53) - let parameters = NWParameters.udp - parameters.prohibitedInterfaceTypes = [.other] - - let conn = NWConnection(to: endpoint, using: parameters) - conn.stateUpdateHandler = { [weak self] state in - Task { - await self?.handleConnectionStateUpdate(state, for: conn) - } + let connection = NWConnection(to: self.dnsServerEndpoint, using: parameters) + connection.stateUpdateHandler = { [weak self] state in + Task { + await self?.handleConnectionStateUpdate(state, for: connection) } - conn.start(queue: .global()) - self.connections.append(conn) } + connection.start(queue: .global()) + self.connection = connection self.cleanupTask = Task { [weak self] in while !Task.isCancelled { @@ -72,7 +85,7 @@ actor DNSLocalClient { } 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 } @@ -84,19 +97,18 @@ actor DNSLocalClient { self.pendingRequests[transactionID] = PendingRequest(tracker: tracker) let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID) - var hasReadyConnection = false - for conn in self.connections where conn.state == .ready { - hasReadyConnection = true - conn.send(content: rewrittenPayload, completion: .contentProcessed({ error in - if let error { - SDLLogger.log("[DNSLocalClient] send error: \(error.localizedDescription)", for: .debug) + connection.send(content: rewrittenPayload, completion: .contentProcessed { error in + if let error { + self.eventContinuation.yield(.sendFailed(error)) + Task { + await self.removePendingRequest(forKey: transactionID) } - })) - } - - if !hasReadyConnection { - self.pendingRequests.removeValue(forKey: transactionID) - } + } + }) + } + + private func removePendingRequest(forKey id: UInt16) { + self.pendingRequests.removeValue(forKey: id) } func stop() { @@ -105,69 +117,52 @@ actor DNSLocalClient { } self.state = .stopped - self.receiveTasks.values.forEach { $0.cancel() } - self.receiveTasks.removeAll() - self.connections.forEach { $0.cancel() } - self.connections.removeAll() + + self.receiveTask?.cancel() + self.receiveTask = nil + + self.connection?.cancel() + self.connection = nil self.cleanupTask?.cancel() self.cleanupTask = nil self.pendingRequests.removeAll() self.nextTransactionID = 1 - self.finishPacketFlowIfNeeded() + + self.packetContinuation.finish() } private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) { - guard case .running = self.state else { - return - } - switch state { case .ready: self.startReceiveTask(for: conn) + self.state = .running case .failed(let error): SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug) - self.stop() + self.eventContinuation.yield(.failed(error)) case .cancelled: - let key = ObjectIdentifier(conn) - self.receiveTasks.removeValue(forKey: key)?.cancel() - self.connections.removeAll { $0 === conn } - if self.connections.isEmpty { - self.stop() - } + self.eventContinuation.yield(.cancelled) default: () } } private func startReceiveTask(for conn: NWConnection) { - let key = ObjectIdentifier(conn) - guard self.receiveTasks[key] == nil else { - return - } - let stream = Self.makeReceiveStream(for: conn) - self.receiveTasks[key] = Task { [weak self] in + + self.receiveTask = Task { [weak self] in for await data in stream { guard let self else { break } 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) { - guard case .running = self.state, - let rewrittenTransactionID = Self.readTransactionID(from: data), + guard let rewrittenTransactionID = Self.readTransactionID(from: data), let pendingRequest = self.pendingRequests.removeValue(forKey: rewrittenTransactionID) else { return } @@ -185,10 +180,6 @@ actor DNSLocalClient { } private func performCleanup() { - guard case .running = self.state else { - return - } - let now = Date() self.pendingRequests = self.pendingRequests.filter { _, request in now.timeIntervalSince(request.tracker.createdAt) < self.timeoutInterval @@ -211,15 +202,6 @@ actor DNSLocalClient { return nil } - private func finishPacketFlowIfNeeded() { - guard !self.didFinishPacketFlow else { - return - } - - self.didFinishPacketFlow = true - self.packetContinuation.finish() - } - private static func nextTransactionID(after id: UInt16) -> UInt16 { return id == UInt16.max ? 1 : id &+ 1 }