diff --git a/Tun/DNS/DNSCloudClient.swift b/Tun/DNS/DNSCloudClient.swift index 0f28d3f..bd5053d 100644 --- a/Tun/DNS/DNSCloudClient.swift +++ b/Tun/DNS/DNSCloudClient.swift @@ -7,29 +7,33 @@ import Foundation import Network -final class DNSCloudClient { +actor DNSCloudClient { enum DNSCloudError: Error { case failed(Error) case cancelled case sendFailed(Error) + case invalidData } private enum State { case idle + case starting case running case stopped } private var state: State = .idle + private let queue = DispatchQueue(label: "com.sdl.DNSCloudClient.queue") private var connection: NWConnection? private let dnsServerAddress: NWEndpoint // 用于对外输出收到的 DNS 响应包 - public let packetFlow: AsyncThrowingStream + nonisolated let packetFlow: AsyncThrowingStream private let packetContinuation: AsyncThrowingStream.Continuation private var isPacketContinuationFinished: Bool = false + private let readySignal = AsyncOneShot() /// - Parameter serverIP: 你的 sn-server IP 地址 (如 "8.8.8.8") /// - Parameter port: 端口 (如 53) @@ -57,6 +61,7 @@ final class DNSCloudClient { guard self.state == .idle else { return } + self.state = .starting // 1. 配置参数:这是解决环路的关键 let parameters = NWParameters.udp @@ -68,17 +73,20 @@ final class DNSCloudClient { // 2. 创建连接 let connection = NWConnection(to: self.dnsServerAddress, using: parameters) connection.stateUpdateHandler = { [weak self] state in - self?.handleConnectionStateUpdate(state, for: connection) + Task { + await self?.handleConnectionStateUpdate(state, for: connection) + } } - // 启动连接队列 - connection.start(queue: .global()) - self.connection = connection - let stream = Self.makeReceiveStream(for: connection) + // 启动连接队列 + connection.start(queue: self.queue) + try await withTaskCancellationHandler { - for await data in stream { + try await self.readySignal.wait() + while true { try Task.checkCancellation() + let data = try await self.readOnce() self.packetContinuation.yield(data) } } onCancel: { @@ -88,45 +96,65 @@ final class DNSCloudClient { /// 发送 DNS 查询包(由 TUN 拦截到的原始 IP 包数据) func forward(ipPacketData: Data) { - guard let connection = self.connection, connection.state == .ready else { + guard self.state == .running, + let connection = self.connection, connection.state == .ready else { return } - connection.send(content: ipPacketData, completion: .contentProcessed { error in - if let error = error { - self.finishPacketContinuationIfNeed(throwing: .sendFailed(error)) + connection.send(content: ipPacketData, completion: .contentProcessed { [weak self] error in + if let error { + Task { + await self?.finishPacketContinuationIfNeed(throwing: .sendFailed(error)) + } } }) } - func stop() { + func stop() async { guard self.state != .stopped else { return } self.state = .stopped - self.connection?.cancel() + let connection = self.connection self.connection = nil + connection?.cancel() + await self.readySignal.fail(DNSCloudError.cancelled) self.finishPacketContinuationIfNeed(throwing: nil) SDLLogger.log("[SDLCloudClient] stopped") } - private func handleConnectionStateUpdate(_ state: NWConnection.State, for connection: NWConnection) { + private func handleConnectionStateUpdate(_ state: NWConnection.State, for connection: NWConnection) async { switch state { case .ready: - SDLLogger.log("[DNSClient] Connection ready", for: .debug) + guard self.state != .stopped, self.isCurrentConnection(connection) else { + return + } + self.state = .running + SDLLogger.log("[DNSClient] Connection ready", for: .debug) + await self.readySignal.succeed(()) case .failed(let error): + await self.readySignal.fail(DNSCloudError.failed(error)) self.finishPacketContinuationIfNeed(throwing: .failed(error)) case .cancelled: + await self.readySignal.fail(DNSCloudError.cancelled) self.finishPacketContinuationIfNeed(throwing: .cancelled) default: break } } + + private func isCurrentConnection(_ connection: NWConnection) -> Bool { + guard let currentConnection = self.connection else { + return false + } + + return currentConnection === connection + } private func finishPacketContinuationIfNeed(throwing error: DNSCloudError?) { guard !self.isPacketContinuationFinished else { @@ -141,25 +169,21 @@ final class DNSCloudClient { } } - /// 接收数据的递归循环 - private static func makeReceiveStream(for connection: NWConnection) -> AsyncStream { - return AsyncStream(bufferingPolicy: .bufferingNewest(256)) { continuation in - func receiveNext() { - connection.receiveMessage { content, _, _, error in - if let data = content, !data.isEmpty { - // 将收到的 DNS 响应写回 AsyncStream - continuation.yield(data) - } - - if error == nil && connection.state == .ready { - receiveNext() // 继续监听下一个包 - } else { - continuation.finish() - } + private func readOnce() async throws -> Data { + guard let connection = self.connection, connection.state == .ready else { + throw DNSCloudError.cancelled + } + + return try await withCheckedThrowingContinuation { continuation in + connection.receiveMessage { content, _, _, error in + if let error { + continuation.resume(throwing: error) + } else if let data = content, !data.isEmpty { + continuation.resume(returning: data) + } else { + continuation.resume(throwing: DNSCloudError.invalidData) } } - - receiveNext() } } diff --git a/Tun/DNS/DNSCloudService.swift b/Tun/DNS/DNSCloudService.swift index 4f7ee21..8632fcf 100644 --- a/Tun/DNS/DNSCloudService.swift +++ b/Tun/DNS/DNSCloudService.swift @@ -23,29 +23,29 @@ actor DNSCloudService { do { try await self.run(client: client) self.clearCurrent(client, generation: generation) - client.stop() + await client.stop() } catch is CancellationError { self.clearCurrent(client, generation: generation) - client.stop() + await client.stop() throw CancellationError() } catch { self.clearCurrent(client, generation: generation) - client.stop() + await client.stop() throw error } } - func stop() { + func stop() async { self.generation &+= 1 let client = self.currentClient self.currentClient = nil - client?.stop() + await client?.stop() } - func forward(ipPacketData: Data) { - self.currentClient?.forward(ipPacketData: ipPacketData) + func forward(ipPacketData: Data) async { + await self.currentClient?.forward(ipPacketData: ipPacketData) } private func run(client: DNSCloudClient) async throws {