调整Context的owns
This commit is contained in:
parent
5b37ac2552
commit
66e2877044
@ -31,10 +31,10 @@ actor ArpServer {
|
||||
return
|
||||
}
|
||||
|
||||
self.cleanupTask = Task {
|
||||
self.cleanupTask = Task { [weak self] in
|
||||
while !Task.isCancelled {
|
||||
try? await Task.sleep(for: .seconds(1))
|
||||
self.cleanup()
|
||||
await self?.cleanup()
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -69,20 +69,24 @@ actor ArpServer {
|
||||
self.known_macs = [:]
|
||||
self.coolingDown = [:]
|
||||
}
|
||||
|
||||
func arpRequest(targetIp: UInt32, use superClient: SDLSuperClient?) async throws {
|
||||
guard let superClient, self.coolingDown[targetIp] == nil else {
|
||||
return
|
||||
|
||||
func stop() {
|
||||
self.cleanupTask?.cancel()
|
||||
self.cleanupTask = nil
|
||||
self.clear()
|
||||
}
|
||||
|
||||
func makeArpRequest(targetIp: UInt32) throws -> Data? {
|
||||
guard self.coolingDown[targetIp] == nil else {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 单位时间内指允许提交一次
|
||||
|
||||
self.coolingDown[targetIp] = Date().addingTimeInterval(3)
|
||||
|
||||
// 进行arp查询
|
||||
|
||||
var arpRequest = SDLArpRequest()
|
||||
arpRequest.targetIp = targetIp
|
||||
|
||||
await superClient.send(type: .arpRequest, data: try arpRequest.serializedData())
|
||||
|
||||
return try arpRequest.serializedData()
|
||||
}
|
||||
|
||||
func handleArpResponse(arpResponse: SDLArpResponse) {
|
||||
|
||||
@ -64,84 +64,61 @@ actor SDLPuncherActor {
|
||||
}
|
||||
}
|
||||
|
||||
func submitRegisterRequest(superClient: SDLSuperClient?, request: RegisterRequest) async {
|
||||
guard let superClient else {
|
||||
return
|
||||
}
|
||||
|
||||
func makeQueryInfoRequest(request: RegisterRequest) async -> Data? {
|
||||
let now = Date()
|
||||
self.cleanupExpiredEntries(now: now)
|
||||
|
||||
|
||||
if let entry = self.requestEntries[request.dstMac], !entry.canSubmit(at: now) {
|
||||
return
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
var queryInfo = SDLQueryInfo()
|
||||
queryInfo.dstMac = request.dstMac
|
||||
|
||||
|
||||
guard let queryData = try? queryInfo.serializedData() else {
|
||||
SDLLogger.log("[SDLPuncherActor] failed to encode queryInfo", for: .debug)
|
||||
return
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
self.requestEntries[request.dstMac] = RequestEntry(
|
||||
request: request,
|
||||
cooldownUntil: now.addingTimeInterval(self.cooldownInterval),
|
||||
phase: .waitingPeerInfo(deadline: now.addingTimeInterval(self.peerInfoTimeout))
|
||||
)
|
||||
|
||||
await superClient.send(type: .queryInfo, data: queryData)
|
||||
|
||||
return queryData
|
||||
}
|
||||
|
||||
func handlePeerInfo(using udpHole: SDLUDPHole?, udpHoleV6: SDLUDPHoleV6?, peerInfo: SDLPeerInfo) async {
|
||||
|
||||
func makeRegisterPackets(peerInfo: SDLPeerInfo) async -> [(data: Data, remoteAddress: SocketAddress)] {
|
||||
let now = Date()
|
||||
self.cleanupExpiredEntries(now: now)
|
||||
|
||||
guard var entry = self.requestEntries[peerInfo.dstMac] else {
|
||||
return
|
||||
|
||||
guard var entry = self.requestEntries[peerInfo.dstMac], entry.isWaitingPeerInfo(at: now) else {
|
||||
return []
|
||||
}
|
||||
|
||||
guard entry.isWaitingPeerInfo(at: now) else {
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
entry.markCoolingDown()
|
||||
self.requestEntries[peerInfo.dstMac] = entry
|
||||
|
||||
guard udpHole != nil || udpHoleV6 != nil else {
|
||||
SDLLogger.log("[SDLPuncherActor] udpHole and udpHoleV6 are nil when peerInfo arrived", for: .debug)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
var register = SDLRegister()
|
||||
register.networkID = entry.request.networkId
|
||||
register.srcMac = entry.request.srcMac
|
||||
register.dstMac = entry.request.dstMac
|
||||
|
||||
|
||||
guard let registerData = try? register.serializedData() else {
|
||||
SDLLogger.log("[SDLPuncherActor] failed to encode register", for: .debug)
|
||||
return
|
||||
return []
|
||||
}
|
||||
|
||||
// 并行发送register请求
|
||||
if peerInfo.hasV4Info {
|
||||
if let remoteAddress = try? await peerInfo.v4Info.socketAddress() {
|
||||
SDLLogger.log("[SDLContext] hole sock address: \(remoteAddress)", for: .debug)
|
||||
await self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress)
|
||||
} else {
|
||||
SDLLogger.log("[SDLPuncherActor] failed to resolve peerInfo.v4Info", for: .debug)
|
||||
}
|
||||
|
||||
var packets: [(data: Data, remoteAddress: SocketAddress)] = []
|
||||
if peerInfo.hasV4Info, let remoteAddress = try? await peerInfo.v4Info.socketAddress() {
|
||||
packets.append((data: registerData, remoteAddress: remoteAddress))
|
||||
}
|
||||
|
||||
if peerInfo.hasV6Info {
|
||||
if let remoteAddress = try? await peerInfo.v6Info.socketAddress() {
|
||||
SDLLogger.log("[SDLContext] hole sock address v6: \(remoteAddress)", for: .debug)
|
||||
await self.sendRegister(using: udpHole, udpHoleV6: udpHoleV6, registerData: registerData, remoteAddress: remoteAddress)
|
||||
} else {
|
||||
SDLLogger.log("[SDLPuncherActor] failed to resolve peerInfo.v6Info", for: .debug)
|
||||
}
|
||||
if peerInfo.hasV6Info, let remoteAddress = try? await peerInfo.v6Info.socketAddress() {
|
||||
packets.append((data: registerData, remoteAddress: remoteAddress))
|
||||
}
|
||||
|
||||
|
||||
return packets
|
||||
}
|
||||
|
||||
func stop() {
|
||||
@ -156,25 +133,6 @@ actor SDLPuncherActor {
|
||||
}
|
||||
}
|
||||
|
||||
private func sendRegister(using udpHole: SDLUDPHole?, udpHoleV6: SDLUDPHoleV6?, registerData: Data, remoteAddress: SocketAddress) async {
|
||||
switch remoteAddress {
|
||||
case .v4:
|
||||
guard let udpHole else {
|
||||
SDLLogger.log("[SDLPuncherActor] udpHole is nil when v4 peerInfo arrived", for: .debug)
|
||||
return
|
||||
}
|
||||
await udpHole.send(type: .register, data: registerData, remoteAddress: remoteAddress)
|
||||
case .v6:
|
||||
guard let udpHoleV6 else {
|
||||
SDLLogger.log("[SDLPuncherActor] udpHoleV6 is nil when v6 peerInfo arrived", for: .debug)
|
||||
return
|
||||
}
|
||||
udpHoleV6.send(type: .register, data: registerData, remoteAddress: remoteAddress)
|
||||
default:
|
||||
SDLLogger.log("[SDLPuncherActor] unsupported peer address family: \(remoteAddress)", for: .debug)
|
||||
}
|
||||
}
|
||||
|
||||
deinit {
|
||||
self.cleanupTask?.cancel()
|
||||
}
|
||||
|
||||
@ -14,7 +14,7 @@ import NIOCore
|
||||
1. 处理rsa的加解密逻辑
|
||||
*/
|
||||
|
||||
private func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task<Void, Never> {
|
||||
func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task<Void, Never> {
|
||||
return Task(name: name) {
|
||||
while true {
|
||||
do {
|
||||
@ -54,20 +54,6 @@ enum SDLContextError: Error {
|
||||
|
||||
actor SDLContextActor {
|
||||
|
||||
private enum UDPHoleKind: Equatable {
|
||||
case v4
|
||||
case v6
|
||||
|
||||
func convertAddressType() -> Session.AddressType {
|
||||
switch self {
|
||||
case .v4:
|
||||
return .v4
|
||||
case .v6:
|
||||
return .v6
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private var config: SDLConfiguration
|
||||
// nat的网络类型
|
||||
var natType: SDLNATProberActor.NatType = .blocked
|
||||
@ -82,26 +68,12 @@ actor SDLContextActor {
|
||||
// 加密算法相关
|
||||
nonisolated let rsaCipher: RSACipher
|
||||
|
||||
// 依赖的变量
|
||||
private var udpHole: SDLUDPHole?
|
||||
private var udpHoleMonitorTask: Task<Void, Never>?
|
||||
private var natProbeTask: Task<Void, Never>?
|
||||
private var udpHoleLocalAddress: SocketAddress?
|
||||
|
||||
private var udpHoleV6: SDLUDPHoleV6?
|
||||
private var udpHoleV6MonitorTask: Task<Void, Never>?
|
||||
private var udpHoleService: SDLUDPHoleService?
|
||||
private var dnsService: SDLDNSService?
|
||||
private var superService: SDLSuperService?
|
||||
private var packetReaderService: SDLPacketReaderService?
|
||||
|
||||
// dns的client对象
|
||||
private var dnsClient: DNSCloudClient?
|
||||
private var dnsMonitorTask: Task<Void, Never>?
|
||||
|
||||
// Localdns的client对象
|
||||
private let publicDnsServers = ["223.5.5.5", "119.29.29.29"]
|
||||
private var dnsLocalClient: DNSLocalClient?
|
||||
private var dnsLocalMonitorTask: Task<Void, Never>?
|
||||
|
||||
private var superClient: SDLSuperClient?
|
||||
private var superMonitorTask: Task<Void, Never>?
|
||||
|
||||
nonisolated private let puncherActor: SDLPuncherActor
|
||||
// 网络探测对象
|
||||
@ -110,9 +82,6 @@ actor SDLContextActor {
|
||||
// 本地ipv6地址信息探测
|
||||
private var ipv6AssistClient: SDLIPV6AssistClient?
|
||||
|
||||
// 数据包读取任务
|
||||
private var readTask: Task<Void, Error>?
|
||||
|
||||
private let sessionManager = SessionManager()
|
||||
nonisolated private let arpServer: ArpServer
|
||||
|
||||
@ -161,76 +130,33 @@ actor SDLContextActor {
|
||||
await self.puncherActor.start()
|
||||
await self.arpServer.start()
|
||||
|
||||
self.startDnsMonitor()
|
||||
self.startDnsLocalMonitor()
|
||||
|
||||
self.startUDPHoleMonitor()
|
||||
|
||||
// self.startUDPHoleV6Monitor()
|
||||
|
||||
self.startSuperMonitor()
|
||||
let dnsService = SDLDNSService(serverHost: self.config.serverHost, publicDnsServers: self.publicDnsServers) { [weak self] event in
|
||||
await self?.handleDNSEvent(event)
|
||||
}
|
||||
self.dnsService = dnsService
|
||||
await dnsService.start()
|
||||
|
||||
let udpHoleService = SDLUDPHoleService(proberActor: self.proberActor) { [weak self] event in
|
||||
await self?.handleUDPHoleEvent(event)
|
||||
}
|
||||
self.udpHoleService = udpHoleService
|
||||
await udpHoleService.start()
|
||||
|
||||
let superService = SDLSuperService(host: self.config.serverHost) { [weak self] message in
|
||||
await self?.handleSuperMessage(message: message)
|
||||
}
|
||||
self.superService = superService
|
||||
await superService.start()
|
||||
}
|
||||
|
||||
// 处理context的停止问题
|
||||
public func stop() async {
|
||||
await self.puncherActor.stop()
|
||||
await self.arpServer.clear()
|
||||
await self.arpServer.stop()
|
||||
await self.sessionManager.clear()
|
||||
|
||||
self.flowSessionManager.clear()
|
||||
|
||||
let udpHole = self.udpHole
|
||||
self.udpHole = nil
|
||||
self.udpHoleLocalAddress = nil
|
||||
|
||||
let udpHoleMonitorTask = self.udpHoleMonitorTask
|
||||
self.udpHoleMonitorTask = nil
|
||||
let natProbeTask = self.natProbeTask
|
||||
self.natProbeTask = nil
|
||||
|
||||
udpHoleMonitorTask?.cancel()
|
||||
natProbeTask?.cancel()
|
||||
await self.proberActor.cancelAll()
|
||||
await udpHole?.stop()
|
||||
if let natProbeTask {
|
||||
await natProbeTask.value
|
||||
}
|
||||
if let udpHoleMonitorTask {
|
||||
await udpHoleMonitorTask.value
|
||||
}
|
||||
|
||||
let udpHoleV6 = self.udpHoleV6
|
||||
self.udpHoleV6 = nil
|
||||
let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask
|
||||
self.udpHoleV6MonitorTask = nil
|
||||
udpHoleV6MonitorTask?.cancel()
|
||||
udpHoleV6?.stop()
|
||||
if let udpHoleV6MonitorTask {
|
||||
await udpHoleV6MonitorTask.value
|
||||
}
|
||||
|
||||
let dnsClient = self.dnsClient
|
||||
self.dnsClient = nil
|
||||
self.dnsMonitorTask?.cancel()
|
||||
self.dnsMonitorTask = nil
|
||||
dnsClient?.stop()
|
||||
|
||||
let dnsLocalClient = self.dnsLocalClient
|
||||
self.dnsLocalClient = nil
|
||||
self.dnsLocalMonitorTask?.cancel()
|
||||
self.dnsLocalMonitorTask = nil
|
||||
await dnsLocalClient?.stop()
|
||||
|
||||
let superClient = self.superClient
|
||||
self.superClient = nil
|
||||
self.superMonitorTask?.cancel()
|
||||
self.superMonitorTask = nil
|
||||
await superClient?.stop()
|
||||
|
||||
SDLLogger.log("[SDLContext] try to cancel readTask")
|
||||
self.readTask?.cancel()
|
||||
self.readTask = nil
|
||||
|
||||
self.registerTask?.cancel()
|
||||
self.registerTask = nil
|
||||
|
||||
@ -240,6 +166,22 @@ actor SDLContextActor {
|
||||
self.updatePolicyTask?.cancel()
|
||||
self.updatePolicyTask = nil
|
||||
|
||||
let packetReaderService = self.packetReaderService
|
||||
self.packetReaderService = nil
|
||||
await packetReaderService?.stop()
|
||||
|
||||
let udpHoleService = self.udpHoleService
|
||||
self.udpHoleService = nil
|
||||
await udpHoleService?.stop()
|
||||
|
||||
let dnsService = self.dnsService
|
||||
self.dnsService = nil
|
||||
await dnsService?.stop()
|
||||
|
||||
let superService = self.superService
|
||||
self.superService = nil
|
||||
await superService?.stop()
|
||||
|
||||
self.sessionToken = nil
|
||||
self.dataCipher = nil
|
||||
self.natType = .blocked
|
||||
@ -249,10 +191,6 @@ actor SDLContextActor {
|
||||
}
|
||||
|
||||
deinit {
|
||||
self.udpHole = nil
|
||||
self.udpHoleLocalAddress = nil
|
||||
self.udpHoleV6 = nil
|
||||
self.dnsClient = nil
|
||||
SDLLogger.log("[SDLContext] deinit", for: .debug)
|
||||
}
|
||||
|
||||
@ -264,16 +202,6 @@ extension SDLContextActor {
|
||||
private func setNatType(natType: SDLNATProberActor.NatType) {
|
||||
self.natType = natType
|
||||
}
|
||||
|
||||
// 探测当前网络的类型
|
||||
private func probeNatType() async {
|
||||
guard let udpHole = self.udpHole else {
|
||||
return
|
||||
}
|
||||
// 开始探测nat的类型
|
||||
self.natType = await self.proberActor.probeNatType(using: udpHole)
|
||||
SDLLogger.log("[SDLContext] nat_type is: \(natType)")
|
||||
}
|
||||
}
|
||||
|
||||
// MARK: Notifier通知机制
|
||||
@ -306,100 +234,13 @@ extension SDLContextActor {
|
||||
}
|
||||
|
||||
private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
|
||||
switch remoteAddress {
|
||||
case .v4:
|
||||
guard let udpHole = self.udpHole else {
|
||||
SDLLogger.log("[SDLContext] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug)
|
||||
return
|
||||
}
|
||||
await udpHole.send(type: type, data: data, remoteAddress: remoteAddress)
|
||||
case .v6:
|
||||
guard let udpHoleV6 = self.udpHoleV6 else {
|
||||
SDLLogger.log("[SDLContext] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug)
|
||||
return
|
||||
}
|
||||
udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress)
|
||||
default:
|
||||
SDLLogger.log("[SDLContext] unsupported socket family: \(remoteAddress)", for: .debug)
|
||||
}
|
||||
await self.udpHoleService?.send(type: type, data: data, remoteAddress: remoteAddress)
|
||||
}
|
||||
}
|
||||
|
||||
// MARK: 处理和Super之间的通讯
|
||||
extension SDLContextActor {
|
||||
|
||||
private func startSuperMonitor() {
|
||||
guard self.superMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.superMonitorTask = startMonitorTask(name: "superMonitorTask") {
|
||||
try await self.startSuperClient()
|
||||
}
|
||||
}
|
||||
|
||||
private func startSuperClient() async throws {
|
||||
let superClient = SDLSuperClient(host: self.config.serverHost, port: 1443)
|
||||
self.superClient = superClient
|
||||
await superClient.start()
|
||||
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
try await runSuperClient(superClient)
|
||||
} onCancel: {
|
||||
SDLLogger.log("[SDLContext] startSuperClient onCancel", for: .debug)
|
||||
Task {
|
||||
await superClient.stop()
|
||||
}
|
||||
}
|
||||
await cleanupSuperClient(superClient)
|
||||
} catch {
|
||||
await cleanupSuperClient(superClient)
|
||||
SDLLogger.log("[SDLContext] startSuperClient catch err: \(error)")
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
private func runSuperClient(_ superClient: SDLSuperClient) async throws {
|
||||
try await Task.sleep(for: .seconds(0.5))
|
||||
try Task.checkCancellation()
|
||||
|
||||
SDLLogger.log("[SDLContext] start super client: \(self.config.serverHost)")
|
||||
|
||||
try await withThrowingTaskGroup(of: Void.self) { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
for try await message in await superClient.messageStream {
|
||||
try Task.checkCancellation()
|
||||
await self.handleSuperMessage(message: message)
|
||||
}
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
while true {
|
||||
try await Task.sleep(for: .seconds(5))
|
||||
try Task.checkCancellation()
|
||||
await superClient.send(type: .ping, data: Data())
|
||||
}
|
||||
}
|
||||
|
||||
_ = try await group.next()
|
||||
}
|
||||
}
|
||||
|
||||
private func cleanupSuperClient(_ superClient: SDLSuperClient) async {
|
||||
await superClient.stop()
|
||||
|
||||
if self.superClient === superClient {
|
||||
self.superClient = nil
|
||||
}
|
||||
|
||||
SDLLogger.log("[SDLContext] cleanupSuperClient")
|
||||
}
|
||||
|
||||
private func handleSuperMessage(message: SDLQUICInboundMessage) async {
|
||||
switch message {
|
||||
case .welcome(let welcome):
|
||||
@ -443,7 +284,10 @@ extension SDLContextActor {
|
||||
self.handleRegisterSuperNak(nakPacket: registerSuperNak)
|
||||
case .peerInfo(let peerInfo):
|
||||
SDLLogger.log("[SDLContext] peer message: \(peerInfo)")
|
||||
await self.puncherActor.handlePeerInfo(using: self.udpHole, udpHoleV6: self.udpHoleV6, peerInfo: 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)
|
||||
}
|
||||
case .event(let event):
|
||||
await self.handleEvent(event: event)
|
||||
case .policyReponse(let policyResponse):
|
||||
@ -489,7 +333,7 @@ extension SDLContextActor {
|
||||
do {
|
||||
try await self.setNetworkSettings(config: self.config, dnsServer: DNSHelper.dnsServer)
|
||||
SDLLogger.log("[SDLContext] setNetworkSettings successed")
|
||||
self.startReader()
|
||||
await self.startPacketReader()
|
||||
// 开启权限的定时更新
|
||||
await self.whenRegistedSuper()
|
||||
} catch let err {
|
||||
@ -507,7 +351,10 @@ extension SDLContextActor {
|
||||
while true {
|
||||
try await Task.sleep(for: .seconds(300))
|
||||
SDLLogger.log("[SDLContext] updatePolicyTask execute")
|
||||
await self.identifyStore.batUpdatePolicy(using: self.superClient, dstIdentityID: self.config.identityId)
|
||||
let requests = await self.identifyStore.makeBatchPolicyRequests(dstIdentityID: self.config.identityId)
|
||||
for request in requests {
|
||||
await self.superService?.send(type: .policyRequest, data: request)
|
||||
}
|
||||
}
|
||||
} catch let err {
|
||||
SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)")
|
||||
@ -573,171 +420,56 @@ extension SDLContextActor {
|
||||
|
||||
if let registerSuperData = try? registerSuper.serializedData() {
|
||||
SDLLogger.log("[SDLContext] will send register super")
|
||||
await self.superClient?.send(type: .registerSuper, data: registerSuperData)
|
||||
await self.superService?.send(type: .registerSuper, data: registerSuperData)
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// MARK: 处理DnsLocal
|
||||
// MARK: DNS service events
|
||||
extension SDLContextActor {
|
||||
|
||||
private func startDnsLocalMonitor() {
|
||||
guard self.dnsLocalMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.dnsLocalMonitorTask = startMonitorTask(name: "dnsLocalMonitorTask") {
|
||||
try await self.startDnsLocalClient()
|
||||
}
|
||||
}
|
||||
|
||||
private func startDnsLocalClient() async throws {
|
||||
let dnsServer = self.publicDnsServers.randomElement() ?? self.publicDnsServers[0]
|
||||
// 启动dns服务
|
||||
let dnsLocalClient = DNSLocalClient(host: dnsServer)
|
||||
await dnsLocalClient.start()
|
||||
SDLLogger.log("[SDLContext] dnsLocalClient started")
|
||||
self.dnsLocalClient = dnsLocalClient
|
||||
|
||||
defer {
|
||||
self.dnsLocalClient = nil
|
||||
}
|
||||
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
// 处理事件流
|
||||
for try await packet in dnsLocalClient.packetFlow {
|
||||
try Task.checkCancellation()
|
||||
// 要想办法构造一个完整的Ip包
|
||||
let nePacket = NEPacket(data: packet, protocolFamily: 2)
|
||||
self.provider.packetFlow.writePacketObjects([nePacket])
|
||||
}
|
||||
} onCancel: {
|
||||
Task {
|
||||
await dnsLocalClient.stop()
|
||||
}
|
||||
}
|
||||
} catch let err {
|
||||
await dnsLocalClient.stop()
|
||||
throw err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// MARK: 处理DnsCloud
|
||||
extension SDLContextActor {
|
||||
|
||||
private func startDnsMonitor() {
|
||||
guard self.dnsMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.dnsMonitorTask = startMonitorTask(name: "dnsMonitorTask") {
|
||||
try await self.startDnsClient()
|
||||
}
|
||||
}
|
||||
|
||||
private func startDnsClient() async throws {
|
||||
// 启动dns服务
|
||||
let dnsClient = DNSCloudClient(host: self.config.serverHost, port: 15353)
|
||||
self.dnsClient = dnsClient
|
||||
dnsClient.start()
|
||||
|
||||
defer {
|
||||
dnsClient.stop()
|
||||
self.dnsClient = nil
|
||||
}
|
||||
|
||||
try await withTaskCancellationHandler {
|
||||
for try await packet in dnsClient.packetFlow {
|
||||
try Task.checkCancellation()
|
||||
|
||||
let nePacket = NEPacket(data: packet, protocolFamily: 2)
|
||||
self.provider.packetFlow.writePacketObjects([nePacket])
|
||||
}
|
||||
} onCancel: {
|
||||
dnsClient.stop()
|
||||
private func handleDNSEvent(_ event: SDLDNSService.Event) async {
|
||||
switch event {
|
||||
case .packet(let packet):
|
||||
let nePacket = NEPacket(data: packet, protocolFamily: 2)
|
||||
self.provider.packetFlow.writePacketObjects([nePacket])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// MARK: 处理从Hole收到的数据
|
||||
extension SDLContextActor {
|
||||
|
||||
private func startUDPHoleMonitor() {
|
||||
guard self.udpHoleMonitorTask == nil else {
|
||||
return
|
||||
private func handleUDPHoleEvent(_ event: SDLUDPHoleService.Event) async {
|
||||
switch event {
|
||||
case .ready(let localAddress):
|
||||
SDLLogger.log("[SDLContext] udpHole ready: \(localAddress)")
|
||||
case .natType(let natType):
|
||||
self.setNatType(natType: natType)
|
||||
SDLLogger.log("[SDLContext] nat_type is: \(natType)")
|
||||
case .packet(let remoteAddress, let message, let source):
|
||||
await self.handleUDPHolePacket(remoteAddress: remoteAddress, message: message, source: source)
|
||||
case .closed(let error):
|
||||
SDLLogger.log("[SDLContext] udpHole closed: \(error)", for: .debug)
|
||||
}
|
||||
|
||||
self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleMonitorTask") {
|
||||
try await self.startUDPHole()
|
||||
}
|
||||
|
||||
private func handleUDPHolePacket(remoteAddress: SocketAddress, message: SDLHoleMessage, source: SDLUDPHoleKind) async {
|
||||
switch message.inboundMessage {
|
||||
case .control(let message):
|
||||
switch message {
|
||||
case .stunReply(_), .stunProbeReply(_):
|
||||
SDLLogger.log("[SDLContext] get a stun reply", for: .debug)
|
||||
case .register(let register):
|
||||
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: source)
|
||||
case .registerAck(let registerAck):
|
||||
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: source)
|
||||
}
|
||||
case .data(let data):
|
||||
try? await self.handleHoleData(data: data)
|
||||
}
|
||||
}
|
||||
|
||||
private func startUDPHole() async throws {
|
||||
// 启动udp服务器
|
||||
let udpHole = try SDLUDPHole()
|
||||
let localAddress = try await udpHole.start()
|
||||
SDLLogger.log("[SDLContext] udpHole started, on address: \(localAddress)")
|
||||
self.udpHole = udpHole
|
||||
self.udpHoleLocalAddress = localAddress
|
||||
|
||||
defer {
|
||||
if self.udpHole === udpHole {
|
||||
self.udpHole = nil
|
||||
self.udpHoleLocalAddress = nil
|
||||
}
|
||||
}
|
||||
|
||||
// 开始探测nat的类型
|
||||
self.natProbeTask?.cancel()
|
||||
let proberActor = self.proberActor
|
||||
self.natProbeTask = Task { [weak self] in
|
||||
SDLLogger.log("[SDLContext] start probeNatType")
|
||||
if Task.isCancelled {
|
||||
return
|
||||
}
|
||||
let natType = await proberActor.probeNatType(using: udpHole)
|
||||
if Task.isCancelled {
|
||||
return
|
||||
}
|
||||
await self?.setNatType(natType: natType)
|
||||
}
|
||||
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
for try await (remoteAddress, message) in await udpHole.messageStream() {
|
||||
try Task.checkCancellation()
|
||||
|
||||
switch message.inboundMessage {
|
||||
case .control(let message):
|
||||
switch message {
|
||||
case .stunReply(_):
|
||||
SDLLogger.log("[SDLContext] get a stunReply", for: .debug)
|
||||
case .stunProbeReply(let probeReply):
|
||||
await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply)
|
||||
case .register(let register):
|
||||
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v4)
|
||||
case .registerAck(let registerAck):
|
||||
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
|
||||
}
|
||||
case .data(let data):
|
||||
try? await self.handleHoleData(data: data)
|
||||
}
|
||||
}
|
||||
} onCancel: {
|
||||
Task {
|
||||
await udpHole.stop()
|
||||
}
|
||||
}
|
||||
} catch let err {
|
||||
await udpHole.stop()
|
||||
throw err
|
||||
}
|
||||
}
|
||||
|
||||
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: UDPHoleKind) async throws {
|
||||
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: SDLUDPHoleKind) async throws {
|
||||
let networkAddr = config.networkAddress
|
||||
SDLLogger.log("[SDLContext] register packet: \(register), network_address: \(networkAddr)")
|
||||
|
||||
@ -761,7 +493,7 @@ extension SDLContextActor {
|
||||
}
|
||||
}
|
||||
|
||||
private func handleRegisterAck(remoteAddress: SocketAddress, registerAck: SDLRegisterAck, source: UDPHoleKind) async {
|
||||
private func handleRegisterAck(remoteAddress: SocketAddress, registerAck: SDLRegisterAck, source: SDLUDPHoleKind) async {
|
||||
// 判断目标地址是否是tun的网卡地址, 并且是在同一个网络下
|
||||
let networkAddr = config.networkAddress
|
||||
if registerAck.dstMac == networkAddr.mac && registerAck.networkID == networkAddr.networkId {
|
||||
@ -805,8 +537,9 @@ extension SDLContextActor {
|
||||
SDLLogger.log("[SDLContext] hole identity: \(identityID), allow, data count: \(packetData.count)", for: .trace)
|
||||
case .requestPolicy(let srcIdentityID):
|
||||
SDLLogger.log("[SDLContext] not found identity: \(srcIdentityID) ruleMap", for: .debug)
|
||||
// 向服务器请求权限逻辑
|
||||
await self.identifyStore.policyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.config.identityId, using: self.superClient)
|
||||
if let queryData = await self.identifyStore.makePolicyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.config.identityId) {
|
||||
await self.superService?.send(type: .policyRequest, data: queryData)
|
||||
}
|
||||
case .none:
|
||||
()
|
||||
}
|
||||
@ -814,83 +547,6 @@ extension SDLContextActor {
|
||||
|
||||
}
|
||||
|
||||
// MARK: 处理从HoleV6收到的数据
|
||||
extension SDLContextActor {
|
||||
|
||||
private func startUDPHoleV6Monitor() {
|
||||
guard self.udpHoleV6MonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.udpHoleV6MonitorTask = startMonitorTask(name: "udpHoleV6MonitorTask") {
|
||||
try await self.startUDPHoleV6()
|
||||
}
|
||||
}
|
||||
|
||||
private func startUDPHoleV6() async throws {
|
||||
// 启动udp服务器
|
||||
let udpHoleV6 = try SDLUDPHoleV6()
|
||||
let localAddress = try udpHoleV6.start()
|
||||
self.udpHoleV6 = udpHoleV6
|
||||
|
||||
if let localAddress {
|
||||
SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)")
|
||||
} else {
|
||||
SDLLogger.log("[SDLContext] udpHoleV6 started, no local address")
|
||||
}
|
||||
|
||||
defer {
|
||||
if self.udpHoleV6 === udpHoleV6 {
|
||||
udpHoleV6.stop()
|
||||
self.udpHoleV6 = nil
|
||||
}
|
||||
}
|
||||
|
||||
try await withThrowingTaskGroup { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
// 处理消息流
|
||||
group.addTask {
|
||||
for await (remoteAddress, message) in udpHoleV6.messageStream {
|
||||
try Task.checkCancellation()
|
||||
|
||||
switch message.inboundMessage {
|
||||
case .control(let message):
|
||||
switch message {
|
||||
case .register(let register):
|
||||
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v6)
|
||||
case .registerAck(let registerAck):
|
||||
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
|
||||
default:
|
||||
()
|
||||
}
|
||||
case .data(let data):
|
||||
try? await self.handleHoleData(data: data)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
for await event in udpHoleV6.eventStream {
|
||||
try Task.checkCancellation()
|
||||
|
||||
switch event {
|
||||
case .ready:
|
||||
SDLLogger.log("[SDLContext] udpHoleV6 ready")
|
||||
case .closed, .errorCaught:
|
||||
throw SDLContextError.udpHoleClosed
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
try await group.next()
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// MARK: 和Stun相关的心跳机制
|
||||
extension SDLContextActor {
|
||||
|
||||
@ -955,32 +611,20 @@ extension SDLContextActor {
|
||||
extension SDLContextActor {
|
||||
|
||||
// 开始读取数据, 用单独的线程处理packetFlow
|
||||
private func startReader() {
|
||||
self.readTask?.cancel()
|
||||
// 开启新的任务
|
||||
let provider = self.provider
|
||||
self.readTask = Task(priority: .high) { [weak self] in
|
||||
try await withTaskCancellationHandler {
|
||||
do {
|
||||
repeat {
|
||||
try Task.checkCancellation()
|
||||
let (packets, numbers) = await provider.packetFlow.readPackets()
|
||||
try Task.checkCancellation()
|
||||
for (data, number) in zip(packets, numbers) where number == 2 {
|
||||
if let ipPacket = IPPacket(data) {
|
||||
await self?.dealTunPacket(packet: ipPacket)
|
||||
}
|
||||
}
|
||||
} while true
|
||||
SDLLogger.log("[SDLContext] readTask finish")
|
||||
} catch let err {
|
||||
SDLLogger.log("[SDLContext] readTask catch error: \(err)")
|
||||
throw err
|
||||
}
|
||||
} onCancel: {
|
||||
SDLLogger.log("[SDLContext] readTask onCancel")
|
||||
private func startPacketReader() async {
|
||||
if self.packetReaderService == nil {
|
||||
self.packetReaderService = SDLPacketReaderService(provider: self.provider) { [weak self] event in
|
||||
await self?.handlePacketReaderEvent(event)
|
||||
}
|
||||
}
|
||||
await self.packetReaderService?.start()
|
||||
}
|
||||
|
||||
private func handlePacketReaderEvent(_ event: SDLPacketReaderService.Event) async {
|
||||
switch event {
|
||||
case .packet(let packet):
|
||||
await self.dealTunPacket(packet: packet)
|
||||
}
|
||||
}
|
||||
|
||||
// 取消出口节点的时候,ip地址为: 0.0.0.0
|
||||
@ -1080,10 +724,10 @@ extension SDLContextActor {
|
||||
self.provider.packetFlow.writePacketObjects([nePacket])
|
||||
case .cloudDNS(let name, let ipPacketData):
|
||||
SDLLogger.log("[SDLContext] get cloud dns request: \(name)")
|
||||
self.dnsClient?.forward(ipPacketData: ipPacketData)
|
||||
await self.dnsService?.forward(ipPacketData: ipPacketData)
|
||||
case .localDNS(let name, let payload, let tracker):
|
||||
SDLLogger.log("[SDLContext] get local dns request: \(name)")
|
||||
await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: payload)
|
||||
await self.dnsService?.queryLocal(tracker: tracker, dnsPayload: payload)
|
||||
case .forwardToNextHop(let ip, let type, let data, let kind):
|
||||
await self.forwardPacketToNextHop(ip: ip, type: type, data: data, kind: kind)
|
||||
case .drop(let reason):
|
||||
@ -1106,11 +750,9 @@ extension SDLContextActor {
|
||||
}
|
||||
else {
|
||||
SDLLogger.log("[SDLContext] dstIp: \(asIpAddress(ip)) arp query not found, broadcast", for: .trace)
|
||||
// // 构造arp广播
|
||||
// let arpReqeust = ARPPacket.arpRequest(senderIP: networkAddr.ip, senderMAC: networkAddr.mac, targetIP: dstIp)
|
||||
// await self.routeLayerPacket(dstMac: ARPPacket.broadcastMac , type: .arp, data: arpReqeust.marshal())
|
||||
|
||||
try? await self.arpServer.arpRequest(targetIp: ip, use: self.superClient)
|
||||
if let arpRequest = try? await self.arpServer.makeArpRequest(targetIp: ip) {
|
||||
await self.superService?.send(type: .arpRequest, data: arpRequest)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@ -1149,7 +791,9 @@ extension SDLContextActor {
|
||||
self.flowTracer.inc(num: payload.count, type: .forward)
|
||||
|
||||
// 尝试打洞
|
||||
await self.puncherActor.submitRegisterRequest(superClient: self.superClient, request: request)
|
||||
if let queryData = await self.puncherActor.makeQueryInfoRequest(request: request) {
|
||||
await self.superService?.send(type: .queryInfo, data: queryData)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
144
Tun/Punchnet/Context/SDLDNSService.swift
Normal file
144
Tun/Punchnet/Context/SDLDNSService.swift
Normal file
@ -0,0 +1,144 @@
|
||||
import Foundation
|
||||
import NetworkExtension
|
||||
|
||||
actor SDLDNSService {
|
||||
enum Event {
|
||||
case packet(Data)
|
||||
}
|
||||
|
||||
typealias EventHandler = @Sendable (Event) async -> Void
|
||||
|
||||
private let serverHost: String
|
||||
private let publicDnsServers: [String]
|
||||
private let onEvent: EventHandler
|
||||
|
||||
private var dnsClient: DNSCloudClient?
|
||||
private var dnsMonitorTask: Task<Void, Never>?
|
||||
|
||||
private var dnsLocalClient: DNSLocalClient?
|
||||
private var dnsLocalMonitorTask: Task<Void, Never>?
|
||||
|
||||
init(serverHost: String, publicDnsServers: [String], onEvent: @escaping EventHandler) {
|
||||
self.serverHost = serverHost
|
||||
self.publicDnsServers = publicDnsServers
|
||||
self.onEvent = onEvent
|
||||
}
|
||||
|
||||
func start() {
|
||||
self.startCloud()
|
||||
self.startLocal()
|
||||
}
|
||||
|
||||
func stop() async {
|
||||
let dnsClient = self.dnsClient
|
||||
self.dnsClient = nil
|
||||
let dnsMonitorTask = self.dnsMonitorTask
|
||||
self.dnsMonitorTask = nil
|
||||
|
||||
let dnsLocalClient = self.dnsLocalClient
|
||||
self.dnsLocalClient = nil
|
||||
let dnsLocalMonitorTask = self.dnsLocalMonitorTask
|
||||
self.dnsLocalMonitorTask = nil
|
||||
|
||||
dnsMonitorTask?.cancel()
|
||||
dnsClient?.stop()
|
||||
if let dnsMonitorTask {
|
||||
await dnsMonitorTask.value
|
||||
}
|
||||
|
||||
dnsLocalMonitorTask?.cancel()
|
||||
await dnsLocalClient?.stop()
|
||||
if let dnsLocalMonitorTask {
|
||||
await dnsLocalMonitorTask.value
|
||||
}
|
||||
}
|
||||
|
||||
func forward(ipPacketData: Data) {
|
||||
self.dnsClient?.forward(ipPacketData: ipPacketData)
|
||||
}
|
||||
|
||||
func queryLocal(tracker: DNSLocalClient.DNSTracker, dnsPayload: Data) async {
|
||||
await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: dnsPayload)
|
||||
}
|
||||
|
||||
private func startCloud() {
|
||||
guard self.dnsMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.dnsMonitorTask = startMonitorTask(name: "dnsServiceCloudMonitor") { [weak self] in
|
||||
guard let self else {
|
||||
throw CancellationError()
|
||||
}
|
||||
try await self.runCloud()
|
||||
}
|
||||
}
|
||||
|
||||
private func startLocal() {
|
||||
guard self.dnsLocalMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.dnsLocalMonitorTask = startMonitorTask(name: "dnsServiceLocalMonitor") { [weak self] in
|
||||
guard let self else {
|
||||
throw CancellationError()
|
||||
}
|
||||
try await self.runLocal()
|
||||
}
|
||||
}
|
||||
|
||||
private func runCloud() async throws {
|
||||
let dnsClient = DNSCloudClient(host: self.serverHost, port: 15353)
|
||||
self.dnsClient = dnsClient
|
||||
dnsClient.start()
|
||||
|
||||
defer {
|
||||
dnsClient.stop()
|
||||
if self.dnsClient === dnsClient {
|
||||
self.dnsClient = nil
|
||||
}
|
||||
}
|
||||
|
||||
let onEvent = self.onEvent
|
||||
try await withTaskCancellationHandler {
|
||||
for try await packet in dnsClient.packetFlow {
|
||||
try Task.checkCancellation()
|
||||
await onEvent(.packet(packet))
|
||||
}
|
||||
} onCancel: {
|
||||
dnsClient.stop()
|
||||
}
|
||||
}
|
||||
|
||||
private func runLocal() async throws {
|
||||
let dnsServer = self.publicDnsServers.randomElement() ?? "223.5.5.5"
|
||||
let dnsLocalClient = DNSLocalClient(host: dnsServer)
|
||||
await dnsLocalClient.start()
|
||||
self.dnsLocalClient = dnsLocalClient
|
||||
SDLLogger.log("[SDLDNSService] dnsLocalClient started")
|
||||
|
||||
defer {
|
||||
if self.dnsLocalClient === dnsLocalClient {
|
||||
self.dnsLocalClient = nil
|
||||
}
|
||||
}
|
||||
|
||||
let onEvent = self.onEvent
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
for try await packet in dnsLocalClient.packetFlow {
|
||||
try Task.checkCancellation()
|
||||
await onEvent(.packet(packet))
|
||||
}
|
||||
} onCancel: {
|
||||
Task {
|
||||
await dnsLocalClient.stop()
|
||||
}
|
||||
}
|
||||
await dnsLocalClient.stop()
|
||||
} catch {
|
||||
await dnsLocalClient.stop()
|
||||
throw error
|
||||
}
|
||||
}
|
||||
}
|
||||
49
Tun/Punchnet/Context/SDLPacketReaderService.swift
Normal file
49
Tun/Punchnet/Context/SDLPacketReaderService.swift
Normal file
@ -0,0 +1,49 @@
|
||||
import Foundation
|
||||
import NetworkExtension
|
||||
|
||||
actor SDLPacketReaderService {
|
||||
enum Event {
|
||||
case packet(IPPacket)
|
||||
}
|
||||
|
||||
typealias EventHandler = @Sendable (Event) async -> Void
|
||||
|
||||
private let provider: NEPacketTunnelProvider
|
||||
private let onEvent: EventHandler
|
||||
private var readTask: Task<Void, Never>?
|
||||
|
||||
init(provider: NEPacketTunnelProvider, onEvent: @escaping EventHandler) {
|
||||
self.provider = provider
|
||||
self.onEvent = onEvent
|
||||
}
|
||||
|
||||
func start() {
|
||||
guard self.readTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
let provider = self.provider
|
||||
let onEvent = self.onEvent
|
||||
self.readTask = Task(priority: .high) {
|
||||
while !Task.isCancelled {
|
||||
let (packets, numbers) = await provider.packetFlow.readPackets()
|
||||
if Task.isCancelled {
|
||||
break
|
||||
}
|
||||
|
||||
for (data, number) in zip(packets, numbers) where number == 2 {
|
||||
if let packet = IPPacket(data) {
|
||||
await onEvent(.packet(packet))
|
||||
}
|
||||
}
|
||||
}
|
||||
SDLLogger.log("[SDLPacketReaderService] readTask finished")
|
||||
}
|
||||
}
|
||||
|
||||
func stop() {
|
||||
let readTask = self.readTask
|
||||
self.readTask = nil
|
||||
readTask?.cancel()
|
||||
}
|
||||
}
|
||||
111
Tun/Punchnet/Context/SDLSuperService.swift
Normal file
111
Tun/Punchnet/Context/SDLSuperService.swift
Normal file
@ -0,0 +1,111 @@
|
||||
import Foundation
|
||||
|
||||
actor SDLSuperService {
|
||||
typealias MessageHandler = @Sendable (SDLQUICInboundMessage) async -> Void
|
||||
|
||||
private let host: String
|
||||
private let port: UInt16
|
||||
private let onMessage: MessageHandler
|
||||
|
||||
private var superClient: SDLSuperClient?
|
||||
private var monitorTask: Task<Void, Never>?
|
||||
|
||||
init(host: String, port: UInt16 = 1443, onMessage: @escaping MessageHandler) {
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.onMessage = onMessage
|
||||
}
|
||||
|
||||
func start() {
|
||||
guard self.monitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.monitorTask = startMonitorTask(name: "superServiceMonitor") { [weak self] in
|
||||
guard let self else {
|
||||
throw CancellationError()
|
||||
}
|
||||
try await self.runOnce()
|
||||
}
|
||||
}
|
||||
|
||||
func stop() async {
|
||||
let monitorTask = self.monitorTask
|
||||
self.monitorTask = nil
|
||||
|
||||
let superClient = self.superClient
|
||||
self.superClient = nil
|
||||
|
||||
monitorTask?.cancel()
|
||||
await superClient?.stop()
|
||||
|
||||
if let monitorTask {
|
||||
await monitorTask.value
|
||||
}
|
||||
}
|
||||
|
||||
func send(type: SDLPacketType, data: Data) async {
|
||||
await self.superClient?.send(type: type, data: data)
|
||||
}
|
||||
|
||||
private func runOnce() async throws {
|
||||
let superClient = SDLSuperClient(host: self.host, port: self.port)
|
||||
self.superClient = superClient
|
||||
await superClient.start()
|
||||
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
try await self.run(superClient)
|
||||
} onCancel: {
|
||||
Task {
|
||||
await superClient.stop()
|
||||
}
|
||||
}
|
||||
await self.cleanup(superClient)
|
||||
} catch {
|
||||
await self.cleanup(superClient)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
private func run(_ superClient: SDLSuperClient) async throws {
|
||||
try await Task.sleep(for: .seconds(0.5))
|
||||
try Task.checkCancellation()
|
||||
|
||||
SDLLogger.log("[SDLSuperService] start super client: \(self.host)")
|
||||
|
||||
try await withThrowingTaskGroup(of: Void.self) { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
let onMessage = self.onMessage
|
||||
group.addTask {
|
||||
for try await message in await superClient.messageStream {
|
||||
try Task.checkCancellation()
|
||||
await onMessage(message)
|
||||
}
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
while true {
|
||||
try await Task.sleep(for: .seconds(5))
|
||||
try Task.checkCancellation()
|
||||
await superClient.send(type: .ping, data: Data())
|
||||
}
|
||||
}
|
||||
|
||||
_ = try await group.next()
|
||||
}
|
||||
}
|
||||
|
||||
private func cleanup(_ superClient: SDLSuperClient) async {
|
||||
await superClient.stop()
|
||||
|
||||
if self.superClient === superClient {
|
||||
self.superClient = nil
|
||||
}
|
||||
|
||||
SDLLogger.log("[SDLSuperService] cleanup")
|
||||
}
|
||||
}
|
||||
240
Tun/Punchnet/Context/SDLUDPHoleService.swift
Normal file
240
Tun/Punchnet/Context/SDLUDPHoleService.swift
Normal file
@ -0,0 +1,240 @@
|
||||
import Foundation
|
||||
import NIOCore
|
||||
|
||||
enum SDLUDPHoleKind: Equatable {
|
||||
case v4
|
||||
case v6
|
||||
|
||||
func convertAddressType() -> Session.AddressType {
|
||||
switch self {
|
||||
case .v4:
|
||||
return .v4
|
||||
case .v6:
|
||||
return .v6
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
actor SDLUDPHoleService {
|
||||
enum Event {
|
||||
case ready(SocketAddress)
|
||||
case natType(SDLNATProberActor.NatType)
|
||||
case packet(SocketAddress, SDLHoleMessage, source: SDLUDPHoleKind)
|
||||
case closed(Error)
|
||||
}
|
||||
|
||||
typealias EventHandler = @Sendable (Event) async -> Void
|
||||
|
||||
private let proberActor: SDLNATProberActor
|
||||
private let onEvent: EventHandler
|
||||
|
||||
private var udpHole: SDLUDPHole?
|
||||
private var udpHoleMonitorTask: Task<Void, Never>?
|
||||
private var natProbeTask: Task<Void, Never>?
|
||||
private var localAddress: SocketAddress?
|
||||
|
||||
private var udpHoleV6: SDLUDPHoleV6?
|
||||
private var udpHoleV6MonitorTask: Task<Void, Never>?
|
||||
|
||||
init(proberActor: SDLNATProberActor, onEvent: @escaping EventHandler) {
|
||||
self.proberActor = proberActor
|
||||
self.onEvent = onEvent
|
||||
}
|
||||
|
||||
func start(includeV6: Bool = false) {
|
||||
self.startV4()
|
||||
if includeV6 {
|
||||
self.startV6()
|
||||
}
|
||||
}
|
||||
|
||||
func stop() async {
|
||||
let udpHole = self.udpHole
|
||||
self.udpHole = nil
|
||||
self.localAddress = nil
|
||||
|
||||
let udpHoleMonitorTask = self.udpHoleMonitorTask
|
||||
self.udpHoleMonitorTask = nil
|
||||
|
||||
let natProbeTask = self.natProbeTask
|
||||
self.natProbeTask = nil
|
||||
|
||||
udpHoleMonitorTask?.cancel()
|
||||
natProbeTask?.cancel()
|
||||
await self.proberActor.cancelAll()
|
||||
await udpHole?.stop()
|
||||
|
||||
if let natProbeTask {
|
||||
await natProbeTask.value
|
||||
}
|
||||
if let udpHoleMonitorTask {
|
||||
await udpHoleMonitorTask.value
|
||||
}
|
||||
|
||||
let udpHoleV6 = self.udpHoleV6
|
||||
self.udpHoleV6 = nil
|
||||
let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask
|
||||
self.udpHoleV6MonitorTask = nil
|
||||
udpHoleV6MonitorTask?.cancel()
|
||||
udpHoleV6?.stop()
|
||||
if let udpHoleV6MonitorTask {
|
||||
await udpHoleV6MonitorTask.value
|
||||
}
|
||||
}
|
||||
|
||||
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
|
||||
switch remoteAddress {
|
||||
case .v4:
|
||||
guard let udpHole else {
|
||||
SDLLogger.log("[SDLUDPHoleService] 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("[SDLUDPHoleService] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug)
|
||||
return
|
||||
}
|
||||
udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress)
|
||||
default:
|
||||
SDLLogger.log("[SDLUDPHoleService] unsupported socket family: \(remoteAddress)", for: .debug)
|
||||
}
|
||||
}
|
||||
|
||||
private func startV4() {
|
||||
guard self.udpHoleMonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleServiceV4Monitor") { [weak self] in
|
||||
guard let self else {
|
||||
throw CancellationError()
|
||||
}
|
||||
try await self.runV4()
|
||||
}
|
||||
}
|
||||
|
||||
private func runV4() async throws {
|
||||
let udpHole = try SDLUDPHole()
|
||||
let localAddress = try await udpHole.start()
|
||||
self.udpHole = udpHole
|
||||
self.localAddress = localAddress
|
||||
SDLLogger.log("[SDLUDPHoleService] udpHole started, on address: \(localAddress)")
|
||||
|
||||
await self.onEvent(.ready(localAddress))
|
||||
self.startNatProbe(using: udpHole)
|
||||
|
||||
defer {
|
||||
if self.udpHole === udpHole {
|
||||
self.udpHole = nil
|
||||
self.localAddress = nil
|
||||
}
|
||||
}
|
||||
|
||||
do {
|
||||
try await withTaskCancellationHandler {
|
||||
for try await (remoteAddress, message) in await udpHole.messageStream() {
|
||||
try Task.checkCancellation()
|
||||
try await self.handleV4Message(remoteAddress: remoteAddress, message: message)
|
||||
}
|
||||
} onCancel: {
|
||||
Task {
|
||||
await udpHole.stop()
|
||||
}
|
||||
}
|
||||
} catch {
|
||||
await udpHole.stop()
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
private func startNatProbe(using udpHole: SDLUDPHole) {
|
||||
self.natProbeTask?.cancel()
|
||||
let proberActor = self.proberActor
|
||||
let onEvent = self.onEvent
|
||||
self.natProbeTask = Task {
|
||||
if Task.isCancelled {
|
||||
return
|
||||
}
|
||||
let natType = await proberActor.probeNatType(using: udpHole)
|
||||
if Task.isCancelled {
|
||||
return
|
||||
}
|
||||
await onEvent(.natType(natType))
|
||||
}
|
||||
}
|
||||
|
||||
private func handleV4Message(remoteAddress: SocketAddress, message: SDLHoleMessage) async throws {
|
||||
switch message.inboundMessage {
|
||||
case .control(let control):
|
||||
switch control {
|
||||
case .stunProbeReply(let probeReply):
|
||||
await self.proberActor.handleProbeReply(localAddress: self.localAddress, reply: probeReply)
|
||||
default:
|
||||
await self.onEvent(.packet(remoteAddress, message, source: .v4))
|
||||
}
|
||||
case .data:
|
||||
await self.onEvent(.packet(remoteAddress, message, source: .v4))
|
||||
}
|
||||
}
|
||||
|
||||
private func startV6() {
|
||||
guard self.udpHoleV6MonitorTask == nil else {
|
||||
return
|
||||
}
|
||||
|
||||
self.udpHoleV6MonitorTask = startMonitorTask(name: "udpHoleServiceV6Monitor") { [weak self] in
|
||||
guard let self else {
|
||||
throw CancellationError()
|
||||
}
|
||||
try await self.runV6()
|
||||
}
|
||||
}
|
||||
|
||||
private func runV6() async throws {
|
||||
let udpHoleV6 = try SDLUDPHoleV6()
|
||||
let localAddress = try udpHoleV6.start()
|
||||
self.udpHoleV6 = udpHoleV6
|
||||
|
||||
if let localAddress {
|
||||
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, on address: \(localAddress)")
|
||||
} else {
|
||||
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, no local address")
|
||||
}
|
||||
|
||||
defer {
|
||||
if self.udpHoleV6 === udpHoleV6 {
|
||||
udpHoleV6.stop()
|
||||
self.udpHoleV6 = nil
|
||||
}
|
||||
}
|
||||
|
||||
try await withThrowingTaskGroup(of: Void.self) { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
let onEvent = self.onEvent
|
||||
group.addTask {
|
||||
for await (remoteAddress, message) in udpHoleV6.messageStream {
|
||||
try Task.checkCancellation()
|
||||
await onEvent(.packet(remoteAddress, message, source: .v6))
|
||||
}
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
for await event in udpHoleV6.eventStream {
|
||||
try Task.checkCancellation()
|
||||
switch event {
|
||||
case .ready:
|
||||
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 ready")
|
||||
case .closed, .errorCaught:
|
||||
throw SDLContextError.udpHoleClosed
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_ = try await group.next()
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -24,49 +24,35 @@ actor IdentityStore {
|
||||
init(publisher: SnapshotPublisher<IdentitySnapshot>) {
|
||||
self.publisher = publisher
|
||||
}
|
||||
|
||||
// 批量更新, 有外部任务驱动,因为这里依赖于当前的superClient
|
||||
func batUpdatePolicy(using superClient: SDLSuperClient?, dstIdentityID: UInt32) async {
|
||||
guard let superClient else {
|
||||
return
|
||||
}
|
||||
|
||||
for identityId in self.identityMap.keys {
|
||||
|
||||
func makeBatchPolicyRequests(dstIdentityID: UInt32) -> [Data] {
|
||||
return self.identityMap.keys.compactMap { identityId in
|
||||
var policyRequest = SDLPolicyRequest()
|
||||
policyRequest.srcIdentityID = identityId
|
||||
policyRequest.dstIdentityID = dstIdentityID
|
||||
policyRequest.version = self.nextVersion(identityId: identityId)
|
||||
|
||||
// 发送请求
|
||||
if let queryData = try? policyRequest.serializedData() {
|
||||
await superClient.send(type: .policyRequest, data: queryData)
|
||||
}
|
||||
return try? policyRequest.serializedData()
|
||||
}
|
||||
}
|
||||
|
||||
// 提交权限请求
|
||||
func policyRequest(srcIdentityId: UInt32, dstIdentityId: UInt32, using superClient: SDLSuperClient?) async {
|
||||
guard let superClient, !coolingDown.contains(srcIdentityId) else {
|
||||
return
|
||||
|
||||
func makePolicyRequest(srcIdentityId: UInt32, dstIdentityId: UInt32) -> Data? {
|
||||
guard !coolingDown.contains(srcIdentityId) else {
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
var policyRequest = SDLPolicyRequest()
|
||||
policyRequest.srcIdentityID = srcIdentityId
|
||||
policyRequest.dstIdentityID = dstIdentityId
|
||||
policyRequest.version = self.nextVersion(identityId: srcIdentityId)
|
||||
|
||||
// 触发一次打洞
|
||||
|
||||
coolingDown.insert(srcIdentityId)
|
||||
// 发送请求
|
||||
if let queryData = try? policyRequest.serializedData() {
|
||||
await superClient.send(type: .policyRequest, data: queryData)
|
||||
}
|
||||
|
||||
Task {
|
||||
// 启动冷却期
|
||||
|
||||
Task { [weak self] in
|
||||
try? await Task.sleep(for: .seconds(5))
|
||||
self.endCooldown(for: srcIdentityId)
|
||||
await self?.endCooldown(for: srcIdentityId)
|
||||
}
|
||||
|
||||
return try? policyRequest.serializedData()
|
||||
}
|
||||
|
||||
// 处理权限的响应
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user