This commit is contained in:
anlicheng 2026-05-27 22:16:02 +08:00
parent 0e455b187d
commit 75ed6d2970
6 changed files with 199 additions and 147 deletions

View File

@ -183,9 +183,13 @@ actor SDLContextActor {
private func runRootBody() async throws { private func runRootBody() async throws {
self.prepareTunnelNotifier() self.prepareTunnelNotifier()
let dnsService = DNSService(serverIP: self.config.serverEndpoint.ip, publicDnsServers: self.publicDnsServers) { [weak self] event in let dnsCloudService = DNSCloudService(serverIP: self.config.serverEndpoint.ip) { [weak self] event in
await self?.handleDNSEvent(event) await self?.handleDNSEvent(event)
} }
let dnsLocalService = DNSLocalService(publicDnsServers: self.publicDnsServers) { [weak self] event in
await self?.handleDNSEvent(event)
}
let dnsService = DNSService(cloudService: dnsCloudService, localService: dnsLocalService)
self.dnsService = dnsService self.dnsService = dnsService
await self.packetOutboundActor.updateDNSService(dnsService) await self.packetOutboundActor.updateDNSService(dnsService)
@ -244,7 +248,15 @@ actor SDLContextActor {
} }
group.addTask { group.addTask {
try await dnsService.run() try await Self.runRestarting(name: "dnsCloudService") {
try await dnsCloudService.run()
}
}
group.addTask {
try await Self.runRestarting(name: "dnsLocalService") {
try await dnsLocalService.run()
}
} }
group.addTask { group.addTask {
@ -637,6 +649,7 @@ extension SDLContextActor {
// MARK: Hole // MARK: Hole
extension SDLContextActor { extension SDLContextActor {
private func handleUDPHoleControlEvent(_ event: SDLUDPHoleService.Event) async { private func handleUDPHoleControlEvent(_ event: SDLUDPHoleService.Event) async {
switch event { switch event {
case .ready(let localAddress): case .ready(let localAddress):

View File

@ -75,10 +75,6 @@ final class DNSCloudClient {
self.connection = connection self.connection = connection
defer {
self.stop()
}
let stream = Self.makeReceiveStream(for: connection) let stream = Self.makeReceiveStream(for: connection)
try await withTaskCancellationHandler { try await withTaskCancellationHandler {
for await data in stream { for await data in stream {
@ -86,7 +82,7 @@ final class DNSCloudClient {
self.packetContinuation.yield(data) self.packetContinuation.yield(data)
} }
} onCancel: { } onCancel: {
self.stop() connection.cancel()
} }
} }

View File

@ -0,0 +1,85 @@
import Foundation
actor DNSCloudService {
private let serverIP: String
private let onEvent: DNSService.EventHandler
private var currentClient: DNSCloudClient?
private var generation: UInt64 = 0
init(serverIP: String, onEvent: @escaping DNSService.EventHandler) {
self.serverIP = serverIP
self.onEvent = onEvent
}
func run() async throws {
let generation = self.nextGeneration()
let client = DNSCloudClient(serverIP: self.serverIP, port: 15353)
self.currentClient = client
do {
try await self.run(client: client)
self.clearCurrent(client, generation: generation)
client.stop()
} catch is CancellationError {
self.clearCurrent(client, generation: generation)
client.stop()
throw CancellationError()
} catch {
self.clearCurrent(client, generation: generation)
client.stop()
throw error
}
}
func stop() {
self.generation &+= 1
let client = self.currentClient
self.currentClient = nil
client?.stop()
}
func forward(ipPacketData: Data) {
self.currentClient?.forward(ipPacketData: ipPacketData)
}
private func run(client: DNSCloudClient) async throws {
let onEvent = self.onEvent
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await client.run()
}
group.addTask {
for try await packet in client.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
}
_ = try await group.next()
}
}
private func nextGeneration() -> UInt64 {
self.generation &+= 1
return self.generation
}
private func clearCurrent(_ client: DNSCloudClient, generation: UInt64) {
guard self.generation == generation else {
return
}
if self.currentClient === client {
self.currentClient = nil
}
}
}

View File

@ -86,10 +86,6 @@ actor DNSLocalClient {
connection.start(queue: .global()) connection.start(queue: .global())
defer {
self.stop()
}
try await withTaskCancellationHandler { try await withTaskCancellationHandler {
try await withThrowingTaskGroup(of: Void.self) { group in try await withThrowingTaskGroup(of: Void.self) { group in
defer { defer {

View File

@ -0,0 +1,88 @@
import Foundation
actor DNSLocalService {
private let publicDnsServers: [String]
private let onEvent: DNSService.EventHandler
private var currentClient: DNSLocalClient?
private var generation: UInt64 = 0
init(publicDnsServers: [String], onEvent: @escaping DNSService.EventHandler) {
self.publicDnsServers = publicDnsServers
self.onEvent = onEvent
}
func run() async throws {
let generation = self.nextGeneration()
let dnsServer = self.publicDnsServers.randomElement() ?? "223.5.5.5"
let client = DNSLocalClient(host: dnsServer)
self.currentClient = client
SDLLogger.log("[DNSLocalService] dnsLocalClient started")
do {
try await self.run(client: client)
self.clearCurrent(client, generation: generation)
await client.stop()
} catch is CancellationError {
self.clearCurrent(client, generation: generation)
await client.stop()
throw CancellationError()
} catch {
self.clearCurrent(client, generation: generation)
await client.stop()
throw error
}
}
func stop() async {
self.generation &+= 1
let client = self.currentClient
self.currentClient = nil
await client?.stop()
}
func query(tracker: DNSLocalClient.DNSTracker, dnsPayload: Data) async {
await self.currentClient?.query(tracker: tracker, dnsPayload: dnsPayload)
}
private func run(client: DNSLocalClient) async throws {
let onEvent = self.onEvent
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await client.run()
}
group.addTask {
for try await packet in client.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
}
_ = try await group.next()
}
}
private func nextGeneration() -> UInt64 {
self.generation &+= 1
return self.generation
}
private func clearCurrent(_ client: DNSLocalClient, generation: UInt64) {
guard self.generation == generation else {
return
}
if self.currentClient === client {
self.currentClient = nil
}
}
}

View File

@ -1,5 +1,4 @@
import Foundation import Foundation
import NetworkExtension
actor DNSService { actor DNSService {
enum Event { enum Event {
@ -8,149 +7,24 @@ actor DNSService {
typealias EventHandler = @Sendable (Event) async -> Void typealias EventHandler = @Sendable (Event) async -> Void
private let serverIP: String private let cloudService: DNSCloudService
private let publicDnsServers: [String] private let localService: DNSLocalService
private let onEvent: EventHandler
private var dnsClient: DNSCloudClient? init(cloudService: DNSCloudService, localService: DNSLocalService) {
private var dnsLocalClient: DNSLocalClient? self.cloudService = cloudService
self.localService = localService
init(serverIP: String, publicDnsServers: [String], onEvent: @escaping EventHandler) {
self.serverIP = serverIP
self.publicDnsServers = publicDnsServers
self.onEvent = onEvent
}
func run() async throws {
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await Self.runRestarting(name: "dnsServiceCloud") {
try await self.runCloud()
}
}
group.addTask {
try await Self.runRestarting(name: "dnsServiceLocal") {
try await self.runLocal()
}
}
try await group.waitForAll()
}
} }
func stop() async { func stop() async {
let dnsClient = self.dnsClient await self.cloudService.stop()
self.dnsClient = nil await self.localService.stop()
let dnsLocalClient = self.dnsLocalClient
self.dnsLocalClient = nil
dnsClient?.stop()
await dnsLocalClient?.stop()
} }
func forward(ipPacketData: Data) { func forward(ipPacketData: Data) async {
self.dnsClient?.forward(ipPacketData: ipPacketData) await self.cloudService.forward(ipPacketData: ipPacketData)
} }
func queryLocal(tracker: DNSLocalClient.DNSTracker, dnsPayload: Data) async { func queryLocal(tracker: DNSLocalClient.DNSTracker, dnsPayload: Data) async {
await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: dnsPayload) await self.localService.query(tracker: tracker, dnsPayload: dnsPayload)
}
private func runCloud() async throws {
let dnsClient = DNSCloudClient(serverIP: self.serverIP, port: 15353)
self.dnsClient = dnsClient
defer {
dnsClient.stop()
if self.dnsClient === dnsClient {
self.dnsClient = nil
}
}
let onEvent = self.onEvent
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await dnsClient.run()
}
group.addTask {
for try await packet in dnsClient.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
}
try await group.next()
}
}
private func runLocal() async throws {
let dnsServer = self.publicDnsServers.randomElement() ?? "223.5.5.5"
let dnsLocalClient = DNSLocalClient(host: dnsServer)
self.dnsLocalClient = dnsLocalClient
SDLLogger.log("[DNSService] dnsLocalClient started")
defer {
if self.dnsLocalClient === dnsLocalClient {
self.dnsLocalClient = nil
}
}
let onEvent = self.onEvent
do {
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await dnsLocalClient.run()
}
group.addTask {
for try await packet in dnsLocalClient.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
}
try await group.next()
}
await dnsLocalClient.stop()
} catch {
await dnsLocalClient.stop()
throw error
}
}
private static func runRestarting(
name: String,
retryDelay: Duration = .seconds(5),
operation: @escaping @Sendable () async throws -> Void
) async throws {
while !Task.isCancelled {
do {
try Task.checkCancellation()
try await operation()
SDLLogger.log("[DNSService] worker \(name) ended, will restart", for: .debug)
} catch is CancellationError {
SDLLogger.log("[DNSService] worker \(name) cancelled", for: .debug)
throw CancellationError()
} catch {
SDLLogger.log("[DNSService] worker \(name) crashed: \(error.localizedDescription), will restart", for: .debug)
}
try await Task.sleep(for: retryDelay)
}
} }
} }