Compare commits

..

No commits in common. "968b86755b0af830b4fbb2770d7d289c6898653d" and "09af58089225f2a77b95060241073645d9e637b4" have entirely different histories.

3 changed files with 140 additions and 139 deletions

View File

@ -312,24 +312,53 @@ actor SDLContextActor {
private func startDnsClient() async throws {
// dns
let dnsClient = DNSCloudClient(host: self.config.serverHost, port: 15353)
self.dnsClient = dnsClient
dnsClient.start()
SDLLogger.log("[SDLContext] dnsClient started")
self.dnsClient = dnsClient
try await withThrowingTaskGroup { group in
defer {
self.dnsClient = nil
dnsClient.stop()
group.cancelAll()
}
try await withTaskCancellationHandler {
for try await packet in dnsClient.packetFlow {
group.addTask {
for await packet in dnsClient.packetFlow {
try Task.checkCancellation()
let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
}
} onCancel: {
dnsClient.stop()
throw SDLContextError.dnsClientClosed
}
group.addTask {
for await event in dnsClient.eventStream {
try Task.checkCancellation()
switch event {
case .failed(let error):
SDLLogger.log("[SDLContext] dnsClient failed with error: \(error)")
throw error
case .cancelled:
SDLLogger.log("[SDLContext] dnsClient cancelled")
throw SDLContextError.dnsClientCancelled
case .sendFailed(let error):
SDLLogger.log("[SDLContext] dnsClient sendFailed with error: \(error)")
throw error
}
}
}
do {
try await group.next()
dnsClient.stop()
self.dnsClient = nil
} catch let err {
dnsClient.stop()
self.dnsClient = nil
throw err
}
}
}
private func startDnsLocalClient() async throws {
@ -340,24 +369,49 @@ actor SDLContextActor {
SDLLogger.log("[SDLContext] dnsLocalClient started")
self.dnsLocalClient = dnsLocalClient
try await withThrowingTaskGroup { group in
defer {
Task {
self.dnsLocalClient = nil
await dnsLocalClient.stop()
}
group.cancelAll()
}
try await withTaskCancellationHandler {
group.addTask {
//
for try await packet in dnsLocalClient.packetFlow {
for await packet in dnsLocalClient.packetFlow {
try Task.checkCancellation()
// Ip
let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
}
} onCancel: {
Task {
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()
await dnsLocalClient.stop()
self.dnsLocalClient = nil
} catch let err {
await dnsLocalClient.stop()
self.dnsLocalClient = nil
throw err
}
}
}

View File

@ -9,7 +9,7 @@ import Network
final class DNSCloudClient {
enum DNSCloudError: Error {
enum Event {
case failed(Error)
case cancelled
case sendFailed(Error)
@ -28,18 +28,25 @@ final class DNSCloudClient {
private let dnsServerAddress: NWEndpoint
// DNS
public let packetFlow: AsyncThrowingStream<Data, Error>
private let packetContinuation: AsyncThrowingStream<Data, Error>.Continuation
private var isPacketContinuationFinished: Bool = false
public let packetFlow: AsyncStream<Data>
private let packetContinuation: AsyncStream<Data>.Continuation
// Connection
public let eventStream: AsyncStream<Event>
private let eventContinuation: AsyncStream<Event>.Continuation
/// - Parameter host: sn-server ( "8.8.8.8")
/// - Parameter port: ( 53)
init(host: String, port: UInt16) {
self.dnsServerAddress = .hostPort(host: NWEndpoint.Host(host), port: NWEndpoint.Port(integerLiteral: port))
let packetPair = AsyncThrowingStream.makeStream(of: Data.self)
let packetPair = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256))
self.packetFlow = packetPair.stream
self.packetContinuation = packetPair.continuation
let eventPair = AsyncStream.makeStream(of: Event.self)
self.eventStream = eventPair.stream
self.eventContinuation = eventPair.continuation
}
func start() {
@ -69,7 +76,7 @@ final class DNSCloudClient {
connection.send(content: ipPacketData, completion: .contentProcessed { error in
if let error = error {
self.finishPacketContinuationIfNeed(throwing: .sendFailed(error))
self.eventContinuation.yield(.sendFailed(error))
}
})
}
@ -87,7 +94,8 @@ final class DNSCloudClient {
self.connection?.cancel()
self.connection = nil
self.finishPacketContinuationIfNeed(throwing: nil)
self.packetContinuation.finish()
self.eventContinuation.finish()
}
private func handleConnectionStateUpdate(_ state: NWConnection.State, for connection: NWConnection) {
@ -97,9 +105,9 @@ final class DNSCloudClient {
self.startReceiveTask(for: connection)
self.state = .running
case .failed(let error):
self.finishPacketContinuationIfNeed(throwing: .failed(error))
self.eventContinuation.yield(.failed(error))
case .cancelled:
self.finishPacketContinuationIfNeed(throwing: .cancelled)
self.eventContinuation.yield(.cancelled)
default:
break
}
@ -121,19 +129,6 @@ final class DNSCloudClient {
}
}
private func finishPacketContinuationIfNeed(throwing error: DNSCloudError?) {
guard !self.isPacketContinuationFinished else {
return
}
self.isPacketContinuationFinished = true
if let error {
self.packetContinuation.finish(throwing: error)
} else {
self.packetContinuation.finish()
}
}
///
private static func makeReceiveStream(for connection: NWConnection) -> AsyncStream<Data> {
return AsyncStream(bufferingPolicy: .bufferingNewest(256)) { continuation in

View File

@ -16,12 +16,11 @@ actor DNSLocalClient {
private enum State {
case idle
case starting
case running
case stopped
}
enum DNSLocalError: Error {
enum Event {
case failed(Error)
case cancelled
case sendFailed(Error)
@ -36,9 +35,16 @@ actor DNSLocalClient {
private var cleanupTask: Task<Void, Never>?
private let timeoutInterval: TimeInterval = 3.0
nonisolated let packetFlow: AsyncThrowingStream<Data, Error>
private let packetContinuation: AsyncThrowingStream<Data, Error>.Continuation
private var isPacketContinuationFinished: Bool = false
public let packetFlow: AsyncStream<Data>
@ObservationIgnored
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 nextTransactionID: UInt16 = 1
@ -46,21 +52,16 @@ actor DNSLocalClient {
init(host: String) {
self.dnsServerEndpoint = .hostPort(host: NWEndpoint.Host(host), port: 53)
let (stream, continuation) = AsyncThrowingStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256))
let (stream, continuation) = AsyncStream.makeStream(of: Data.self, bufferingPolicy: .bufferingNewest(256))
self.packetFlow = stream
self.packetContinuation = continuation
self.packetContinuation.onTermination = { termination in
SDLLogger.log("[DNSLocalClient] packetFlow terminated: \(termination)")
}
let eventPair = AsyncStream.makeStream(of: Event.self)
self.eventStream = eventPair.stream
self.eventContinuation = eventPair.continuation
}
func start() {
guard self.state == .idle else {
return
}
self.state = .starting
let parameters = NWParameters.udp
parameters.prohibitedInterfaceTypes = [.other]
// 2. pathSelectionOptions
@ -72,43 +73,44 @@ actor DNSLocalClient {
await self?.handleConnectionStateUpdate(state, for: connection)
}
}
connection.start(queue: .global())
self.connection = connection
let cleanupTask = Task { [weak self] in
self.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 self.state != .stopped,
let connection = self.connection, connection.state == .ready, dnsPayload.count >= 2 else {
guard let connection = self.connection, connection.state == .ready, dnsPayload.count >= 2 else {
return
}
guard let allocatedTransactionID = self.allocateTransactionID() else {
guard let transactionID = self.allocateTransactionID() else {
SDLLogger.log("[DNSLocalClient] no available transaction id", for: .debug)
return
}
let transactionID = allocatedTransactionID
self.pendingRequests[transactionID] = PendingRequest(tracker: tracker)
let rewrittenPayload = Self.rewriteTransactionID(in: dnsPayload, to: transactionID)
connection.send(content: rewrittenPayload, completion: .contentProcessed { [weak self] error in
connection.send(content: rewrittenPayload, completion: .contentProcessed { error in
if let error {
self.eventContinuation.yield(.sendFailed(error))
Task {
await self?.handleSendFailure(transactionID: transactionID, error: error)
await self.removePendingRequest(forKey: transactionID)
}
}
})
}
private func removePendingRequest(forKey id: UInt16) {
self.pendingRequests.removeValue(forKey: id)
}
func stop() {
guard self.state != .stopped else {
return
@ -116,35 +118,31 @@ actor DNSLocalClient {
self.state = .stopped
let receiveTask = self.receiveTask
self.receiveTask?.cancel()
self.receiveTask = nil
let connection = self.connection
self.connection?.cancel()
self.connection = nil
let cleanupTask = self.cleanupTask
self.cleanupTask?.cancel()
self.cleanupTask = nil
self.pendingRequests.removeAll()
self.nextTransactionID = 1
receiveTask?.cancel()
connection?.cancel()
cleanupTask?.cancel()
self.finishPacketContinuationIfNeed(throwing: nil)
self.packetContinuation.finish()
}
private func handleConnectionStateUpdate(_ state: NWConnection.State, for conn: NWConnection) {
switch state {
case .ready:
if self.markConnectionReady(conn) {
self.startReceiveTask(for: conn)
}
self.state = .running
case .failed(let error):
SDLLogger.log("[DNSLocalClient] failed with error: \(error.localizedDescription)", for: .debug)
self.finishPacketContinuationIfNeed(throwing: .failed(error))
self.eventContinuation.yield(.failed(error))
case .cancelled:
self.finishPacketContinuationIfNeed(throwing: .cancelled)
self.eventContinuation.yield(.cancelled)
default:
()
}
@ -153,7 +151,7 @@ actor DNSLocalClient {
private func startReceiveTask(for conn: NWConnection) {
let stream = Self.makeReceiveStream(for: conn)
let task = Task { [weak self] in
self.receiveTask = Task { [weak self] in
for await data in stream {
guard let self else {
break
@ -161,30 +159,6 @@ 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) {
@ -205,28 +179,6 @@ 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()
self.pendingRequests = self.pendingRequests.filter { _, request in