diff --git a/Tun/Punchnet/Actors/ArpServer.swift b/Tun/Punchnet/Actors/ArpServer.swift index 5607ac6..ef27215 100644 --- a/Tun/Punchnet/Actors/ArpServer.swift +++ b/Tun/Punchnet/Actors/ArpServer.swift @@ -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) { diff --git a/Tun/Punchnet/Actors/SDLPuncherActor.swift b/Tun/Punchnet/Actors/SDLPuncherActor.swift index 4f1bc4c..7ac801c 100644 --- a/Tun/Punchnet/Actors/SDLPuncherActor.swift +++ b/Tun/Punchnet/Actors/SDLPuncherActor.swift @@ -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() } diff --git a/Tun/Punchnet/Context/SDLContextActor.swift b/Tun/Punchnet/Context/SDLContextActor.swift index 6e7b339..461d79f 100644 --- a/Tun/Punchnet/Context/SDLContextActor.swift +++ b/Tun/Punchnet/Context/SDLContextActor.swift @@ -14,7 +14,7 @@ import NIOCore 1. 处理rsa的加解密逻辑 */ -private func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task { +func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task { 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? - private var natProbeTask: Task? - private var udpHoleLocalAddress: SocketAddress? - - private var udpHoleV6: SDLUDPHoleV6? - private var udpHoleV6MonitorTask: Task? + 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? - - // Localdns的client对象 private let publicDnsServers = ["223.5.5.5", "119.29.29.29"] - private var dnsLocalClient: DNSLocalClient? - private var dnsLocalMonitorTask: Task? - - private var superClient: SDLSuperClient? - private var superMonitorTask: Task? nonisolated private let puncherActor: SDLPuncherActor // 网络探测对象 @@ -110,9 +82,6 @@ actor SDLContextActor { // 本地ipv6地址信息探测 private var ipv6AssistClient: SDLIPV6AssistClient? - // 数据包读取任务 - private var readTask: Task? - 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) + } } } diff --git a/Tun/Punchnet/Context/SDLDNSService.swift b/Tun/Punchnet/Context/SDLDNSService.swift new file mode 100644 index 0000000..6244a1e --- /dev/null +++ b/Tun/Punchnet/Context/SDLDNSService.swift @@ -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? + + private var dnsLocalClient: DNSLocalClient? + private var dnsLocalMonitorTask: Task? + + 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 + } + } +} diff --git a/Tun/Punchnet/Context/SDLPacketReaderService.swift b/Tun/Punchnet/Context/SDLPacketReaderService.swift new file mode 100644 index 0000000..581cf22 --- /dev/null +++ b/Tun/Punchnet/Context/SDLPacketReaderService.swift @@ -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? + + 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() + } +} diff --git a/Tun/Punchnet/Context/SDLSuperService.swift b/Tun/Punchnet/Context/SDLSuperService.swift new file mode 100644 index 0000000..7974c46 --- /dev/null +++ b/Tun/Punchnet/Context/SDLSuperService.swift @@ -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? + + 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") + } +} diff --git a/Tun/Punchnet/Context/SDLUDPHoleService.swift b/Tun/Punchnet/Context/SDLUDPHoleService.swift new file mode 100644 index 0000000..d7a40bb --- /dev/null +++ b/Tun/Punchnet/Context/SDLUDPHoleService.swift @@ -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? + private var natProbeTask: Task? + private var localAddress: SocketAddress? + + private var udpHoleV6: SDLUDPHoleV6? + private var udpHoleV6MonitorTask: Task? + + 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() + } + } +} diff --git a/Tun/Punchnet/Policy/IdentityStore.swift b/Tun/Punchnet/Policy/IdentityStore.swift index e5e9d3a..a088198 100644 --- a/Tun/Punchnet/Policy/IdentityStore.swift +++ b/Tun/Punchnet/Policy/IdentityStore.swift @@ -24,49 +24,35 @@ actor IdentityStore { init(publisher: SnapshotPublisher) { 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() } // 处理权限的响应