diff --git a/Tun/Context/SDLContextActor.swift b/Tun/Context/SDLContextActor.swift index 24c735f..0e22286 100644 --- a/Tun/Context/SDLContextActor.swift +++ b/Tun/Context/SDLContextActor.swift @@ -43,6 +43,7 @@ actor SDLContextActor { private var dnsService: DNSService? private let superService: SDLSuperService private let udpHoleService: SDLUDPHoleService + private let udpHoleV6Service: SDLUDPHoleV6Service private let packetOutboundActor: PacketOutboundActor private let packetInboundActor: PacketInboundActor private let tunNetworkManager: SDLTunNetworkManager @@ -83,6 +84,7 @@ actor SDLContextActor { let policyService = PolicyService(identityId: config.identityId, acl: config.acl) let superService = SDLSuperService(serverEndpoint: config.serverEndpoint) let udpHoleService = SDLUDPHoleService(proberActor: proberActor) + let udpHoleV6Service = SDLUDPHoleV6Service() let tunNetworkManager = SDLTunNetworkManager(provider: provider) let ipv6AssistPair = AsyncStream.makeStream(of: Optional.self, bufferingPolicy: .bufferingNewest(1)) let packetOutboundActor = PacketOutboundActor( @@ -95,6 +97,7 @@ actor SDLContextActor { policyService: policyService, superService: superService, udpHoleService: udpHoleService, + udpHoleV6Service: udpHoleV6Service, flowTracer: flowTracer ) let packetInboundActor = PacketInboundActor( @@ -123,6 +126,7 @@ actor SDLContextActor { self.policyService = policyService self.superService = superService self.udpHoleService = udpHoleService + self.udpHoleV6Service = udpHoleV6Service self.packetOutboundActor = packetOutboundActor self.packetInboundActor = packetInboundActor self.tunNetworkManager = tunNetworkManager @@ -198,9 +202,18 @@ actor SDLContextActor { await packetInboundActor.handleData(data) } ) + await self.udpHoleV6Service.updateHandlers( + onEvent: { [weak self] event in + await self?.handleUDPHoleControlEvent(event) + }, + onData: { data in + await packetInboundActor.handleData(data) + } + ) let superService = self.superService let udpHoleService = self.udpHoleService + let udpHoleV6Service = self.udpHoleV6Service let packetOutboundActor = self.packetOutboundActor let policyService = self.policyService let puncherActor = self.puncherActor @@ -220,7 +233,13 @@ actor SDLContextActor { group.addTask { try await Self.runRestarting(name: "udpHoleService") { - try await udpHoleService.run(includeV6: false) + try await udpHoleService.run() + } + } + + group.addTask { + try await Self.runRestarting(name: "udpHoleV6Service") { + try await udpHoleV6Service.run() } } @@ -357,6 +376,7 @@ actor SDLContextActor { await self.policyService.clear() await self.udpHoleService.stop() + await self.udpHoleV6Service.stop() let dnsService = self.dnsService self.dnsService = nil @@ -419,7 +439,14 @@ extension SDLContextActor { } private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { - await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) + switch remoteAddress { + case .v4: + await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) + case .v6: + await self.udpHoleV6Service.send(type: type, data: data, remoteAddress: remoteAddress) + default: + SDLLogger.log("[SDLContext] unsupported socket family: \(remoteAddress)", for: .debug) + } } } @@ -453,7 +480,7 @@ extension SDLContextActor { SDLLogger.log("[SDLContext] peer message: \(peerInfo)") let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo) for packet in packets { - await self.udpHoleService.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress) + await self.sendPacket(type: .register, data: packet.data, remoteAddress: packet.remoteAddress) } case .event(let event): await self.handleEvent(event: event) diff --git a/Tun/Outbound/PacketOutboundActor.swift b/Tun/Outbound/PacketOutboundActor.swift index fdd28c8..53f4f46 100644 --- a/Tun/Outbound/PacketOutboundActor.swift +++ b/Tun/Outbound/PacketOutboundActor.swift @@ -25,6 +25,7 @@ actor PacketOutboundActor { private let policyService: PolicyService private let superService: SDLSuperService private let udpHoleService: SDLUDPHoleService + private let udpHoleV6Service: SDLUDPHoleV6Service private let flowTracer: SDLFlowTracer private var networkAddress: SDLConfiguration.NetworkAddress @@ -43,6 +44,7 @@ actor PacketOutboundActor { policyService: PolicyService, superService: SDLSuperService, udpHoleService: SDLUDPHoleService, + udpHoleV6Service: SDLUDPHoleV6Service, flowTracer: SDLFlowTracer) { self.provider = provider self.networkAddress = config.networkAddress @@ -56,6 +58,7 @@ actor PacketOutboundActor { self.policyService = policyService self.superService = superService self.udpHoleService = udpHoleService + self.udpHoleV6Service = udpHoleV6Service self.flowTracer = flowTracer } @@ -233,6 +236,13 @@ actor PacketOutboundActor { } private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { - await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) + switch remoteAddress { + case .v4: + await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress) + case .v6: + await self.udpHoleV6Service.send(type: type, data: data, remoteAddress: remoteAddress) + default: + SDLLogger.log("[PacketOutboundActor] unsupported socket family: \(remoteAddress)", for: .debug) + } } } diff --git a/Tun/UDPHole/SDLNATProberActor.swift b/Tun/UDPHole/SDLNATProberActor.swift index bdc68f0..31932a2 100644 --- a/Tun/UDPHole/SDLNATProberActor.swift +++ b/Tun/UDPHole/SDLNATProberActor.swift @@ -175,19 +175,19 @@ actor SDLNATProberActor { guard !Task.isCancelled else { return } - await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 1, attr: .none), remoteAddress: addressArray[0][0]) + udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 1, attr: .none), remoteAddress: addressArray[0][0]) guard !Task.isCancelled else { return } - await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 2, attr: .none), remoteAddress: addressArray[1][1]) + udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 2, attr: .none), remoteAddress: addressArray[1][1]) guard !Task.isCancelled else { return } - await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 3, attr: .peer), remoteAddress: addressArray[0][0]) + udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 3, attr: .peer), remoteAddress: addressArray[0][0]) guard !Task.isCancelled else { return } - await udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 4, attr: .port), remoteAddress: addressArray[0][0]) + udpHole.send(type: .stunProbe, data: makeProbePacket(cookieId: cookie, step: 4, attr: .port), remoteAddress: addressArray[0][0]) } private func makeProbePacket(cookieId: UInt32, step: UInt32, attr: SDLProbeAttr) -> Data { diff --git a/Tun/UDPHole/SDLUDPHole.swift b/Tun/UDPHole/SDLUDPHole.swift index f160d71..d80a8c4 100644 --- a/Tun/UDPHole/SDLUDPHole.swift +++ b/Tun/UDPHole/SDLUDPHole.swift @@ -7,52 +7,129 @@ import Foundation import NIOCore import NIOPosix -import SwiftProtobuf -actor SDLUDPHole { +// 处理和sn-server服务器之间的通讯 +final class SDLUDPHole: ChannelInboundHandler { + typealias InboundIn = AddressedEnvelope + + struct SDLHoleDatagram { + let remoteAddress: SocketAddress + let message: SDLHoleMessage + } + enum State { case idle case running case stopped } - + private var state: State = .idle - private let udpHoleHandler: SDLUDPHoleHandler - + private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) + private var channel: Channel? + + let messageStream: AsyncThrowingStream + private let messageContinuation: AsyncThrowingStream.Continuation + init() throws { - self.udpHoleHandler = try SDLUDPHoleHandler() + let (stream, continuation) = AsyncThrowingStream.makeStream(of: SDLHoleDatagram.self, bufferingPolicy: .bufferingNewest(2048)) + self.messageStream = stream + self.messageContinuation = continuation } - - func start() async throws -> SocketAddress { - let localAddress = try self.udpHoleHandler.start() + + func start() throws -> SocketAddress { + guard self.state == .idle else { + guard let localAddress = self.channel?.localAddress else { + throw SDLUDPHoleError.invalidLocalAddress + } + + return localAddress + } + + let bootstrap = DatagramBootstrap(group: group) + .channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) + .channelInitializer { channel in + channel.pipeline.addHandler(self) + } + + // 绑定到IPv4通配地址,只处理IPv4流量 + let channel = try bootstrap.bind(host: "0.0.0.0", port: 0).wait() + guard let localAddress = channel.localAddress else { + throw SDLUDPHoleError.invalidLocalAddress + } + + self.channel = channel self.state = .running - + return localAddress } - - func messageStream() -> AsyncThrowingStream { - return self.udpHoleHandler.messageStream + + // --MARK: ChannelInboundHandler delegate + + func channelRead(context: ChannelHandlerContext, data: NIOAny) { + let envelope = unwrapInboundIn(data) + var buffer = envelope.data + let remoteAddress = envelope.remoteAddress + + do { + if let message = try SDLHoleMessage.decode(buffer: &buffer) { + self.messageContinuation.yield(SDLHoleDatagram(remoteAddress: remoteAddress, message: message)) + } else { + SDLLogger.log("[SDLUDPHole] decode message, get null", for: .debug) + } + } catch let err { + SDLLogger.log("[SDLUDPHole] decode message, get error: \(err)", for: .debug) + self.messageContinuation.finish(throwing: err) + } } - + + func channelInactive(context: ChannelHandlerContext) { + self.messageContinuation.finish(throwing: SDLUDPHoleError.closed) + } + + func errorCaught(context: ChannelHandlerContext, error: any Error) { + context.close(promise: nil) + self.messageContinuation.finish(throwing: SDLUDPHoleError.errorCaught) + } + + // MARK: 处理写入逻辑 func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { - guard self.state == .running else { + guard self.state == .running, let channel = self.channel else { return } - - self.udpHoleHandler.send(type: type, data: data, remoteAddress: remoteAddress) + + var buffer = channel.allocator.buffer(capacity: data.count + 1) + buffer.writeBytes([type.rawValue]) + buffer.writeBytes(data) + + let envelope = AddressedEnvelope(remoteAddress: remoteAddress, data: buffer) + let promise = channel.eventLoop.makePromise(of: Void.self) + + channel.eventLoop.execute { + channel.writeAndFlush(envelope, promise: promise) + } + + promise.futureResult.whenFailure { [weak self] err in + self?.messageContinuation.finish(throwing: SDLUDPHoleError.sendFaied(err)) + } } - - func stop() async { + + func stop() { guard self.state != .stopped else { return } - + self.state = .stopped - self.udpHoleHandler.stop() + self.messageContinuation.finish() + + let channel = self.channel + self.channel = nil + try? channel?.close().wait() + try? self.group.syncShutdownGracefully() + + SDLLogger.log("[SDLUDPHole] stopped", for: .debug) } - + deinit { - SDLLogger.log("[SDLUDPHoleActor] deinit", for: .debug) + SDLLogger.log("[SDLUDPHole] deinit", for: .debug) } } - diff --git a/Tun/UDPHole/SDLUDPHoleHandler.swift b/Tun/UDPHole/SDLUDPHoleHandler.swift deleted file mode 100644 index 0d138d2..0000000 --- a/Tun/UDPHole/SDLUDPHoleHandler.swift +++ /dev/null @@ -1,115 +0,0 @@ -// -// SDLUDPHoleHandler.swift -// punchnet -// 上游的SDLUDPHole是基于actor实现的,串行实现的逻辑已经保证了SDLUDPHoleHandler的内部的一致性 -// Created by 安礼成 on 2026/5/22. -// -import Foundation -import NIOCore -import NIOPosix - -// 处理和sn-server服务器之间的通讯 -final class SDLUDPHoleHandler: ChannelInboundHandler { - typealias InboundIn = AddressedEnvelope - - struct SDLHoleDatagram { - let remoteAddress: SocketAddress - let message: SDLHoleMessage - } - - private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1) - private var channel: Channel? - - public let messageStream: AsyncThrowingStream - private let messageContinuation: AsyncThrowingStream.Continuation - - // 启动函数 - init() throws { - let (stream, continuation) = AsyncThrowingStream.makeStream(of: SDLHoleDatagram.self, bufferingPolicy: .bufferingNewest(2048)) - self.messageStream = stream - self.messageContinuation = continuation - } - - func start() throws -> SocketAddress { - let bootstrap = DatagramBootstrap(group: group) - .channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1) - .channelInitializer { channel in - channel.pipeline.addHandler(self) - } - - // 绑定到IPv4通配地址,只处理IPv4流量 - let channel = try bootstrap.bind(host: "0.0.0.0", port: 0).wait() - guard let localAddress = channel.localAddress else { - throw SDLUDPHoleError.invalidLocalAddress - } - - self.channel = channel - - return localAddress - } - - // --MARK: ChannelInboundHandler delegate - - func channelRead(context: ChannelHandlerContext, data: NIOAny) { - let envelope = unwrapInboundIn(data) - var buffer = envelope.data - let remoteAddress = envelope.remoteAddress - - do { - if let message = try SDLHoleMessage.decode(buffer: &buffer) { - self.messageContinuation.yield(SDLHoleDatagram(remoteAddress: remoteAddress, message: message)) - } else { - SDLLogger.log("[SDLUDPHole] decode message, get null", for: .debug) - } - } catch let err { - SDLLogger.log("[SDLUDPHole] decode message, get error: \(err)", for: .debug) - self.messageContinuation.finish(throwing: err) - } - } - - func channelInactive(context: ChannelHandlerContext) { - self.messageContinuation.finish(throwing: SDLUDPHoleError.closed) - } - - func errorCaught(context: ChannelHandlerContext, error: any Error) { - context.close(promise: nil) - self.messageContinuation.finish(throwing: SDLUDPHoleError.errorCaught) - } - - // MARK: 处理写入逻辑 - func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { - guard let channel = self.channel else { - return - } - - var buffer = channel.allocator.buffer(capacity: data.count + 1) - buffer.writeBytes([type.rawValue]) - buffer.writeBytes(data) - - let envelope = AddressedEnvelope(remoteAddress: remoteAddress, data: buffer) - - let promise = channel.eventLoop.makePromise(of: Void.self) - channel.eventLoop.execute { - channel.writeAndFlush(envelope, promise: promise) - } - - promise.futureResult.whenFailure { [weak self] err in - self?.messageContinuation.finish(throwing: SDLUDPHoleError.sendFaied(err)) - } - } - - func stop() { - self.messageContinuation.finish() - let channel = self.channel - self.channel = nil - try? channel?.close().wait() - try? self.group.syncShutdownGracefully() - - SDLLogger.log("[SDLUDPHole] stopped", for: .debug) - } - - deinit { - SDLLogger.log("[SDLUDPHole] deinit", for: .debug) - } - -} diff --git a/Tun/UDPHole/SDLUDPHoleService.swift b/Tun/UDPHole/SDLUDPHoleService.swift index f2036d9..4c45951 100644 --- a/Tun/UDPHole/SDLUDPHoleService.swift +++ b/Tun/UDPHole/SDLUDPHoleService.swift @@ -42,11 +42,10 @@ actor SDLUDPHoleService { self.onData = onData } - func run(includeV6: Bool = false) async throws { + func run() async throws { let generation = self.nextGeneration() let session = SDLUDPHoleSession( proberActor: self.proberActor, - includeV6: includeV6, onEvent: { [weak self] event in await self?.handleEvent(event, generation: generation) }, diff --git a/Tun/UDPHole/SDLUDPHoleSession.swift b/Tun/UDPHole/SDLUDPHoleSession.swift index de375c3..eac83bb 100644 --- a/Tun/UDPHole/SDLUDPHoleSession.swift +++ b/Tun/UDPHole/SDLUDPHoleSession.swift @@ -9,45 +9,25 @@ import NIOCore actor SDLUDPHoleSession { private let proberActor: SDLNATProberActor - private let includeV6: Bool private let onEvent: SDLUDPHoleService.EventHandler private let onData: SDLUDPHoleService.DataHandler private var udpHole: SDLUDPHole? - private var udpHoleV6: SDLUDPHoleV6? private var localAddress: SocketAddress? init( proberActor: SDLNATProberActor, - includeV6: Bool, onEvent: @escaping SDLUDPHoleService.EventHandler, onData: @escaping SDLUDPHoleService.DataHandler ) { self.proberActor = proberActor - self.includeV6 = includeV6 self.onEvent = onEvent self.onData = onData } func run() async throws { do { - try await withThrowingTaskGroup(of: Void.self) { group in - defer { - group.cancelAll() - } - - group.addTask { - try await self.runV4() - } - - if self.includeV6 { - group.addTask { - try await self.runV6() - } - } - - try await group.waitForAll() - } + try await self.runV4() await self.stop() } catch { await self.stop() @@ -59,39 +39,28 @@ actor SDLUDPHoleSession { let udpHole = self.udpHole self.udpHole = nil self.localAddress = nil - await udpHole?.stop() + udpHole?.stop() - let udpHoleV6 = self.udpHoleV6 - self.udpHoleV6 = nil - udpHoleV6?.stop() - await self.proberActor.cancelAll() } func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { - switch remoteAddress { - case .v4: - guard let udpHole else { - SDLLogger.log("[SDLUDPHoleSession] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) - return - } - - await udpHole.send(type: type, data: data, remoteAddress: remoteAddress) - case .v6: - guard let udpHoleV6 else { - SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) - return - } - - udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress) - default: + guard case .v4 = remoteAddress else { SDLLogger.log("[SDLUDPHoleSession] unsupported socket family: \(remoteAddress)", for: .debug) + return } + + guard let udpHole else { + SDLLogger.log("[SDLUDPHoleSession] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) + return + } + + udpHole.send(type: type, data: data, remoteAddress: remoteAddress) } private func runV4() async throws { let udpHole = try SDLUDPHole() - let localAddress = try await udpHole.start() + let localAddress = try udpHole.start() self.udpHole = udpHole self.localAddress = localAddress @@ -115,7 +84,7 @@ actor SDLUDPHoleSession { try await group.waitForAll() } } catch { - await udpHole.stop() + udpHole.stop() if self.udpHole === udpHole { self.udpHole = nil self.localAddress = nil @@ -125,7 +94,7 @@ actor SDLUDPHoleSession { } private func readV4Loop(udpHole: SDLUDPHole) async throws { - for try await datagram in await udpHole.messageStream() { + for try await datagram in udpHole.messageStream { try Task.checkCancellation() try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message) } @@ -157,58 +126,4 @@ actor SDLUDPHoleSession { await self.onData(data) } } - - private func runV6() async throws { - let udpHoleV6 = try SDLUDPHoleV6() - let localAddress = try udpHoleV6.start() - self.udpHoleV6 = udpHoleV6 - - if let localAddress { - SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, on address: \(localAddress)") - } else { - SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, no local address") - } - - do { - try await withThrowingTaskGroup(of: Void.self) { group in - defer { - group.cancelAll() - } - - let onEvent = self.onEvent - let onData = self.onData - group.addTask { - for await (remoteAddress, message) in udpHoleV6.messageStream { - try Task.checkCancellation() - switch message { - case .control(let control): - await onEvent(.packet(remoteAddress, control, source: .v6)) - case .data(let data): - await onData(data) - } - } - } - - group.addTask { - for await event in udpHoleV6.eventStream { - try Task.checkCancellation() - switch event { - case .ready: - SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 ready") - case .closed, .errorCaught: - throw SDLContextError.udpHoleClosed - } - } - } - - _ = try await group.next() - } - } catch { - udpHoleV6.stop() - if self.udpHoleV6 === udpHoleV6 { - self.udpHoleV6 = nil - } - throw error - } - } } diff --git a/Tun/UDPHole/SDLHoleMessage.swift b/Tun/UDPHoleCommon/SDLHoleMessage.swift similarity index 100% rename from Tun/UDPHole/SDLHoleMessage.swift rename to Tun/UDPHoleCommon/SDLHoleMessage.swift diff --git a/Tun/UDPHole/SDLUDPHoleError.swift b/Tun/UDPHoleCommon/SDLUDPHoleError.swift similarity index 100% rename from Tun/UDPHole/SDLUDPHoleError.swift rename to Tun/UDPHoleCommon/SDLUDPHoleError.swift diff --git a/Tun/UDPHole/SDLUDPHoleV6.swift b/Tun/UDPHoleV6/SDLUDPHoleV6.swift similarity index 100% rename from Tun/UDPHole/SDLUDPHoleV6.swift rename to Tun/UDPHoleV6/SDLUDPHoleV6.swift diff --git a/Tun/UDPHoleV6/SDLUDPHoleV6Service.swift b/Tun/UDPHoleV6/SDLUDPHoleV6Service.swift new file mode 100644 index 0000000..26187b4 --- /dev/null +++ b/Tun/UDPHoleV6/SDLUDPHoleV6Service.swift @@ -0,0 +1,78 @@ +import Foundation +import NIOCore + +actor SDLUDPHoleV6Service { + typealias EventHandler = SDLUDPHoleService.EventHandler + typealias DataHandler = SDLUDPHoleService.DataHandler + + private var onEvent: EventHandler = { _ in } + private var onData: DataHandler = { _ in } + private var currentSession: SDLUDPHoleV6Session? + private var generation: UInt64 = 0 + + func updateHandlers(onEvent: @escaping EventHandler, onData: @escaping DataHandler) { + self.onEvent = onEvent + self.onData = onData + } + + func run() async throws { + let generation = self.nextGeneration() + let session = SDLUDPHoleV6Session( + onEvent: { [weak self] event in + await self?.handleEvent(event, generation: generation) + }, + onData: self.onData + ) + + self.currentSession = session + + do { + try await session.run() + self.clearCurrent(session, generation: generation) + } catch is CancellationError { + self.clearCurrent(session, generation: generation) + await session.stop() + throw CancellationError() + } catch { + self.clearCurrent(session, generation: generation) + await session.stop() + throw error + } + } + + func stop() async { + self.generation &+= 1 + + let session = self.currentSession + self.currentSession = nil + + await session?.stop() + } + + func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { + await self.currentSession?.send(type: type, data: data, remoteAddress: remoteAddress) + } + + private func nextGeneration() -> UInt64 { + self.generation &+= 1 + return self.generation + } + + private func clearCurrent(_ session: SDLUDPHoleV6Session, generation: UInt64) { + guard self.generation == generation else { + return + } + + if self.currentSession === session { + self.currentSession = nil + } + } + + private func handleEvent(_ event: SDLUDPHoleService.Event, generation: UInt64) async { + guard self.generation == generation else { + return + } + + await self.onEvent(event) + } +} diff --git a/Tun/UDPHoleV6/SDLUDPHoleV6Session.swift b/Tun/UDPHoleV6/SDLUDPHoleV6Session.swift new file mode 100644 index 0000000..eaeef15 --- /dev/null +++ b/Tun/UDPHoleV6/SDLUDPHoleV6Session.swift @@ -0,0 +1,91 @@ +import Foundation +import NIOCore + +actor SDLUDPHoleV6Session { + private let onEvent: SDLUDPHoleService.EventHandler + private let onData: SDLUDPHoleService.DataHandler + + private var udpHoleV6: SDLUDPHoleV6? + + init( + onEvent: @escaping SDLUDPHoleService.EventHandler, + onData: @escaping SDLUDPHoleService.DataHandler + ) { + self.onEvent = onEvent + self.onData = onData + } + + func run() async throws { + let udpHoleV6 = try SDLUDPHoleV6() + let localAddress = try udpHoleV6.start() + self.udpHoleV6 = udpHoleV6 + + if let localAddress { + SDLLogger.log("[SDLUDPHoleV6Session] udpHoleV6 started, on address: \(localAddress)") + } else { + SDLLogger.log("[SDLUDPHoleV6Session] udpHoleV6 started, no local address") + } + + do { + try await withThrowingTaskGroup(of: Void.self) { group in + defer { + group.cancelAll() + } + + let onEvent = self.onEvent + let onData = self.onData + + group.addTask { + for await (remoteAddress, message) in udpHoleV6.messageStream { + try Task.checkCancellation() + switch message { + case .control(let control): + await onEvent(.packet(remoteAddress, control, source: .v6)) + case .data(let data): + await onData(data) + } + } + } + + group.addTask { + for await event in udpHoleV6.eventStream { + try Task.checkCancellation() + switch event { + case .ready: + SDLLogger.log("[SDLUDPHoleV6Session] udpHoleV6 ready") + case .closed, .errorCaught: + throw SDLContextError.udpHoleClosed + } + } + } + + _ = try await group.next() + } + + await self.stop() + } catch { + await self.stop() + throw error + } + } + + func stop() async { + let udpHoleV6 = self.udpHoleV6 + self.udpHoleV6 = nil + udpHoleV6?.stop() + } + + func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) { + guard case .v6 = remoteAddress else { + SDLLogger.log("[SDLUDPHoleV6Session] unsupported socket family: \(remoteAddress)", for: .debug) + return + } + + guard let udpHoleV6 else { + SDLLogger.log("[SDLUDPHoleV6Session] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) + return + } + + udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress) + } +}