punchnet-macos/Tun/Punchnet/Context/SDLContextActor.swift
2026-05-07 08:43:48 +08:00

1157 lines
42 KiB
Swift
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

//
// SDLContext.swift
// Tun
//
// Created by on 2024/2/29.
//
import Foundation
import NetworkExtension
import NIOCore
//
/*
1. rsa
*/
private func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task<Void, Never> {
return Task(name: name) {
while true {
do {
try Task.checkCancellation()
try await body()
} catch is CancellationError {
SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug)
break
} catch let err {
SDLLogger.log("[SDLContext] worker \(name) crashed: \(err.localizedDescription), will restart", for: .debug)
do {
try await Task.sleep(for: retryDelay)
} catch is CancellationError {
break
} catch {
break
}
}
}
}
}
// ip
private func asIpAddress(_ ipNum: UInt32) -> String {
return SDLUtil.int32ToIp(ipNum)
}
enum SDLContextError: Error {
case udpHoleClosed
case dnsLocalClientClosed
case dnsLocalClientCancelled
case dnsClientClosed
case dnsClientCancelled
}
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
// AES
private var dataCipher: CCDataCipher?
// session token
private var sessionToken: Data?
// rsa, public_key
//
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>?
// dnsclient
private var dnsClient: DNSCloudClient?
private var dnsMonitorTask: Task<Void, Never>?
// Localdnsclient
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
//
nonisolated private let proberActor: SDLNATProberActor
// ipv6
private var ipv6AssistClient: SDLIPV6AssistClient?
//
private var readTask: Task<Void, Error>?
private let sessionManager = SessionManager()
nonisolated private let arpServer: ArpServer
// socket
// App Group + Darwin Notification
//
nonisolated private let flowTracer = SDLFlowTracer()
nonisolated private let provider: NEPacketTunnelProvider
//
private let identifyStore: IdentityStore
private var updatePolicyTask: Task<Void, Never>?
private let snapshotPublisher: SnapshotPublisher<IdentitySnapshot>
// Flow : 180
private let flowSessionManager = SDLFlowSessionManager(sessionTimeout: 180)
//
private var registerTask: Task<Void, Never>?
// stunRequest
private var stunRequestTask: Task<Void, Never>?
public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) {
self.provider = provider
self.config = config
self.rsaCipher = rsaCipher
self.puncherActor = SDLPuncherActor()
self.proberActor = SDLNATProberActor(addressArray: config.stunProbeSocketAddressArray)
self.arpServer = ArpServer()
//
let snapshotPublisher = SnapshotPublisher(initial: IdentitySnapshot.empty())
self.identifyStore = IdentityStore(publisher: snapshotPublisher)
self.snapshotPublisher = snapshotPublisher
}
public func start() async {
self.prepareTunnelNotifier()
// arp
await self.puncherActor.start()
await self.arpServer.start()
self.startDnsMonitor()
self.startDnsLocalMonitor()
self.startUDPHoleMonitor()
// self.startUDPHoleV6Monitor()
self.startSuperMonitor()
}
// context
public func stop() async {
await self.puncherActor.stop()
await self.arpServer.clear()
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
self.stunRequestTask?.cancel()
self.stunRequestTask = nil
self.updatePolicyTask?.cancel()
self.updatePolicyTask = nil
self.sessionToken = nil
self.dataCipher = nil
self.natType = .blocked
await self.ipv6AssistClient?.stop()
self.ipv6AssistClient = nil
}
deinit {
self.udpHole = nil
self.udpHoleLocalAddress = nil
self.udpHoleV6 = nil
self.dnsClient = nil
SDLLogger.log("[SDLContext] deinit", for: .debug)
}
}
// MARK: probe
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
extension SDLContextActor {
private func prepareTunnelNotifier() {
// noticeClient
// UDP NoticeClient App Group
SDLTunnelAppNotifier.shared.clear()
SDLLogger.log("[SDLContext] tunnelAppNotifier ready")
}
private func publishTunnelEvent(code: Int? = nil, message: String) {
SDLTunnelAppNotifier.shared.publish(code: code, message: message)
}
}
// MARK:
extension SDLContextActor {
// super/stun
private func sendSuperPacket(type: SDLPacketType, data: Data) async {
await self.sendPacket(type: type, data: data, remoteAddress: self.config.stunSocketAddress)
}
// peer
private func sendPeerPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
await self.sendPacket(type: type, data: data, remoteAddress: remoteAddress)
}
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)
}
}
}
// 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):
SDLLogger.log("[SDLContext] quic welcome: \(welcome)")
//
await self.doRegisterSuper()
SDLLogger.log("[SDLContext] quic doRegisterSuper")
//
self.registerTask = Task {
do {
try await Task.sleep(for: .seconds(5))
try Task.checkCancellation()
// Tunnel
self.publishTunnelEvent(message: "校验失败")
// 退
let error = NSError(domain: "com.jihe.punchnet.tun", code: -3)
self.provider.cancelTunnelWithError(error)
} catch {
SDLLogger.log("[SDLContext] registerTask: cancel")
return
}
}
// stun
await self.startStunRequestTask(welcome: welcome)
case .pong:
//SDLLogger.shared.log("[SDLContext] quic pong")
()
case .registerSuperAck(let registerSuperAck):
self.registerTask?.cancel()
self.registerTask = nil
await self.handleRegisterSuperAck(registerSuperAck: registerSuperAck)
case .registerSuperNak(let registerSuperNak):
self.registerTask?.cancel()
self.registerTask = nil
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)
case .event(let event):
await self.handleEvent(event: event)
case .policyReponse(let policyResponse):
//
await self.identifyStore.applyPolicyResponse(policyResponse)
case .arpResponse(let arpResponse):
SDLLogger.log("[SDLContext] get arp response: \(arpResponse)")
await self.arpServer.handleArpResponse(arpResponse: arpResponse)
}
}
private func makeSuperEventProcessor() -> SDLSuperEventProcessor {
return .init(networkAddress: self.config.networkAddress)
}
private func handleRegisterSuperAck(registerSuperAck: SDLRegisterSuperAck) async {
// rsa
guard let key = try? self.rsaCipher.decode(data: Data(registerSuperAck.key)) else {
SDLLogger.log("[SDLContext] registerSuperAck invalid key")
let error = SDLError.invalidKey
self.provider.cancelTunnelWithError(error)
return
}
let algorithm = registerSuperAck.algorithm.lowercased()
let regionId = registerSuperAck.regionID
self.sessionToken = registerSuperAck.sessionToken
switch algorithm {
case "aes":
self.dataCipher = CCAESChiper(key: key)
case "chacha20":
self.dataCipher = CCChaCha20Cipher(regionId: regionId, keyData: key)
default:
SDLLogger.log("[SDLContext] registerSuperAck invalid algorithm \(algorithm)")
let error = SDLError.unsupportedAlgorithm(algorithm: algorithm)
self.provider.cancelTunnelWithError(error)
return
}
SDLLogger.log("[SDLContext] registerSuperAck, use algorithm \(algorithm), key len: \(key.count)")
// tun
do {
try await self.setNetworkSettings(config: self.config, dnsServer: DNSHelper.dnsServer)
SDLLogger.log("[SDLContext] setNetworkSettings successed")
self.startReader()
//
await self.whenRegistedSuper()
} catch let err {
SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)")
self.provider.cancelTunnelWithError(err)
}
}
// super
private func whenRegistedSuper() async {
self.updatePolicyTask?.cancel()
self.updatePolicyTask = Task {
do {
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)
}
} catch let err {
SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)")
}
}
}
private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) {
let errorMessage = nakPacket.errorMessage
guard let errorCode = SDLNAKErrorCode(rawValue: UInt8(nakPacket.errorCode)) else {
return
}
switch errorCode {
case .invalidToken, .nodeDisabled:
self.publishTunnelEvent(code: Int(errorCode.rawValue), message: errorMessage)
// 退
let error = NSError(domain: "com.jihe.punchnet.tun", code: -1)
self.provider.cancelTunnelWithError(error)
case .noIpAddress, .networkFault, .internalFault:
self.publishTunnelEvent(code: Int(errorCode.rawValue), message: errorMessage)
}
SDLLogger.log("[SDLContext] Get a SuperNak message exit")
}
private func handleEvent(event: SDLEvent) async {
let processor = self.makeSuperEventProcessor()
let plan = await processor.makeProcessingPlan(event: event)
if let logMessage = plan.logMessage {
SDLLogger.log(logMessage)
}
switch plan.action {
case .removeSession(let dstMac):
await self.sessionManager.removeSession(dstMac: dstMac)
case .sendRegister(let registerData, let remoteAddresses):
for remoteAddress in remoteAddresses {
await self.sendPeerPacket(type: .register, data: registerData, remoteAddress: remoteAddress)
}
case .shutdown(let message):
self.publishTunnelEvent(message: message)
// 退
let error = NSError(domain: "com.jihe.punchnet.tun", code: -2)
self.provider.cancelTunnelWithError(error)
case .none:
()
}
}
private func doRegisterSuper() async {
//
var registerSuper = SDLRegisterSuper()
registerSuper.clientID = self.config.clientId
registerSuper.networkID = self.config.networkAddress.networkId
registerSuper.mac = self.config.networkAddress.mac
registerSuper.ip = self.config.networkAddress.ip
registerSuper.maskLen = UInt32(self.config.networkAddress.maskLen)
registerSuper.hostname = self.config.hostname
registerSuper.pubKey = self.rsaCipher.pubKey
registerSuper.accessToken = self.config.accessToken
if let registerSuperData = try? registerSuper.serializedData() {
SDLLogger.log("[SDLContext] will send register super")
await self.superClient?.send(type: .registerSuper, data: registerSuperData)
}
}
}
// MARK: DnsLocal
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()
}
}
}
// MARK: Hole
extension SDLContextActor {
private func startUDPHoleMonitor() {
guard self.udpHoleMonitorTask == nil else {
return
}
self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleMonitorTask") {
try await self.startUDPHole()
}
}
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 {
let networkAddr = config.networkAddress
SDLLogger.log("[SDLContext] register packet: \(register), network_address: \(networkAddr)")
// tun,
if register.dstMac == networkAddr.mac && register.networkID == networkAddr.networkId {
// ack
var registerAck = SDLRegisterAck()
registerAck.networkID = networkAddr.networkId
registerAck.srcMac = networkAddr.mac
registerAck.dstMac = register.srcMac
await self.sendPeerPacket(type: .registerAck, data: try registerAck.serializedData(), remoteAddress: remoteAddress)
// , super-nodenatudpnat
if let session = Session(dstMac: register.srcMac, natAddress: remoteAddress, addressType: source.convertAddressType()) {
await self.sessionManager.addSession(session: session)
} else {
SDLLogger.log("[SDLContext] didReadRegister get unsupported remoteAddress: \(remoteAddress)", for: .debug)
}
} else {
SDLLogger.log("[SDLContext] didReadRegister get a invalid packet, because dst_ip not matched: \(register.dstMac)")
}
}
private func handleRegisterAck(remoteAddress: SocketAddress, registerAck: SDLRegisterAck, source: UDPHoleKind) async {
// tun,
let networkAddr = config.networkAddress
if registerAck.dstMac == networkAddr.mac && registerAck.networkID == networkAddr.networkId {
if let session = Session(dstMac: registerAck.srcMac, natAddress: remoteAddress, addressType: source.convertAddressType()) {
await self.sessionManager.addSession(session: session)
} else {
SDLLogger.log("[SDLContext] didReadRegisterAck get unsupported remoteAddress: \(remoteAddress)", for: .debug)
}
} else {
SDLLogger.log("[SDLContext] didReadRegisterAck get a invalid packet, because dst_mac not matched: \(registerAck.dstMac)")
}
}
private func makeHoleDataProcessor() -> SDLHoleDataProcessor {
return .init(
networkAddress: self.config.networkAddress,
dataCipher: self.dataCipher,
snapshotPublisher: self.snapshotPublisher,
flowSessionManager: self.flowSessionManager
)
}
private func handleHoleData(data: SDLData) async throws {
let processor = self.makeHoleDataProcessor()
guard let plan = try processor.makeProcessingPlan(data: data) else {
return
}
self.flowTracer.inc(num: plan.inboundBytes, type: .inbound)
switch plan.action {
case .sendARPReply(let dstMac, let responseData):
SDLLogger.log("[SDLContext] get arp request packet")
await self.routeLayerPacket(dstMac: dstMac, type: .arp, data: responseData)
case .appendARP(let ip, let mac):
SDLLogger.log("[SDLContext] get arp response packet")
await self.arpServer.append(ip: ip, mac: mac)
case .writeToTun(let packetData, let identityID):
let packet = NEPacket(data: packetData, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([packet])
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)
case .none:
()
}
}
}
// 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 {
// MARK: -- StunRequestTask
private func startStunRequestTask(welcome: SDLWelcome) async {
self.stunRequestTask?.cancel()
self.stunRequestTask = nil
await self.ipv6AssistClient?.stop()
self.ipv6AssistClient = SDLIPV6AssistClient(assistServerInfo: welcome.ipv6Assist)
await self.ipv6AssistClient?.start()
let ipv6AssistClient = self.ipv6AssistClient
// welcome使ipv6
//
self.stunRequestTask = Task { [weak self] in
let timerStream = SDLAsyncTimerStream()
timerStream.start(interval: .seconds(8))
for await _ in timerStream.stream {
if Task.isCancelled {
break
}
let probeReply = try? await ipv6AssistClient?.probe(requestTimeout: .seconds(3))
if let v6Info = probeReply?.v6Info, let v6Address = SDLUtil.ipv6DataToString(v6Info.v6) {
SDLLogger.log("[SDLContext] probe ipv6 address: \(v6Address)")
} else {
SDLLogger.log("[SDLContext] probe ipv6 address: empty")
}
await self?.sendStunRequest(probeReply: probeReply)
}
SDLLogger.log("[SDLContext] udp stunRequestTask cancel")
}
}
private func sendStunRequest(probeReply: SDLV6AssistProbeReply?) async {
guard let sessionToken = self.sessionToken else {
return
}
var stunRequest = SDLStunRequest()
stunRequest.clientID = self.config.clientId
stunRequest.networkID = self.config.networkAddress.networkId
stunRequest.ip = self.config.networkAddress.ip
stunRequest.mac = self.config.networkAddress.mac
stunRequest.natType = UInt32(self.natType.rawValue)
stunRequest.sessionToken = sessionToken
if let v6Info = probeReply?.v6Info {
stunRequest.v6Info = v6Info
}
if let stunData = try? stunRequest.serializedData() {
await self.sendSuperPacket(type: .stunRequest, data: stunData)
}
}
}
// MARK: NEPacketTunnelProvider
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")
}
}
}
// ip: 0.0.0.0
public func updateExitNode(exitNodeIp: String) async throws {
if let ip = SDLUtil.ipv4StrToInt32(exitNodeIp), ip > 0 {
self.config.exitNode = .init(exitNodeIp: ip)
} else {
self.config.exitNode = nil
}
try await self.setNetworkSettings(config: config, dnsServer: DNSHelper.dnsServer)
}
// MARK:
private func setNetworkSettings(config: SDLConfiguration, dnsServer: String) async throws {
let networkAddress = config.networkAddress
//
var routes: [NEIPv4Route] = [
NEIPv4Route(destinationAddress: networkAddress.netAddress, subnetMask: networkAddress.maskAddress),
NEIPv4Route(destinationAddress: dnsServer, subnetMask: "255.255.255.255"),
]
//
if config.exitNode != nil {
routes.append(.default())
}
// Add code here to start the process of connecting the tunnel.
let networkSettings = NEPacketTunnelNetworkSettings(tunnelRemoteAddress: "8.8.8.8")
networkSettings.mtu = 1250
// DNS
let networkDomain = networkAddress.networkDomain
let dnsSettings = NEDNSSettings(servers: [dnsServer])
dnsSettings.searchDomains = [networkDomain]
dnsSettings.matchDomains = [networkDomain, ""]
// false Search Domain
dnsSettings.matchDomainsNoSearch = false
networkSettings.dnsSettings = dnsSettings
let ipv4Settings = NEIPv4Settings(addresses: [networkAddress.ipAddress], subnetMasks: [networkAddress.maskAddress])
//
ipv4Settings.includedRoutes = routes
//
ipv4Settings.excludedRoutes = self.getIpv4ExcludeRoutes()
networkSettings.ipv4Settings = ipv4Settings
//
try await self.provider.setTunnelNetworkSettings(networkSettings)
}
private func getIpv4ExcludeRoutes() -> [NEIPv4Route] {
//
let dnsServers = SDLUtil.getMacOSSystemDnsServers()
var ipv4DnsServers = dnsServers.filter {!$0.contains(":")}
// dns
let commonDnsServers = [
"8.8.8.8",
"8.8.4.4",
"223.5.5.5",
"223.6.6.6",
"114.114.114.114"
]
for ip in commonDnsServers {
if !ipv4DnsServers.contains(ip) {
ipv4DnsServers.append(ip)
}
}
return ipv4DnsServers.map { NEIPv4Route(destinationAddress: $0, subnetMask: "255.255.255.255") }
}
// , Tun
private func dealTunPacket(packet: IPPacket) async {
let router = SDLTunPacketRouter(networkAddress: self.config.networkAddress, exitNode: self.config.exitNode)
let decision = router.route(packet: packet)
// FlowSession
//
if decision.shouldTrackFlow, let flowSession = packet.flowSession() {
self.flowSessionManager.updateSession(flowSession)
//SDLLogger.shared.log("[SDLContext] flow_session: \(flowSession)", level: .debug)
}
await self.handleTunRouteDecision(decision)
}
private func handleTunRouteDecision(_ decision: SDLTunPacketRouter.RouteDecision) async {
switch decision {
case .loopback(let ipPacketData):
let nePacket = NEPacket(data: ipPacketData, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
case .cloudDNS(let name, let ipPacketData):
SDLLogger.log("[SDLContext] get cloud dns request: \(name)")
self.dnsClient?.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)
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):
SDLLogger.log("[SDLContext] drop tun packet, reason: \(reason.rawValue)", for: .trace)
}
}
private func forwardPacketToNextHop(ip: UInt32, type: LayerPacket.PacketType, data: Data, kind: SDLTunPacketRouter.ForwardKind) async {
switch kind {
case .sameNetwork:
SDLLogger.log("[SDLContext] dstIp: \(asIpAddress(ip)) same network", for: .trace)
case .exitNode, .dnsExitNode:
SDLLogger.log("[SDLContext] use exit_node: \(asIpAddress(ip))", for: .trace)
}
// arpmac
if let dstMac = await self.arpServer.query(ip: ip) {
SDLLogger.log("[SDLContext] dstIp: \(asIpAddress(ip)), dst_mac is: \(SDLUtil.formatMacAddress(mac: dstMac))", for: .trace)
await self.routeLayerPacket(dstMac: dstMac, type: type, data: data)
}
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)
}
}
private func makeLayerPacketForwarder() -> SDLLayerPacketForwarder {
return .init(
networkAddress: self.config.networkAddress,
identityID: self.config.identityId,
dataCipher: self.dataCipher,
sessionManager: self.sessionManager
)
}
private func routeLayerPacket(dstMac: Data, type: LayerPacket.PacketType, data: Data) async {
// 2
//
let forwarder = self.makeLayerPacketForwarder()
guard let plan = try? await forwarder.makeDeliveryPlan(dstMac: dstMac, type: type, data: data) else {
return
}
// 广
switch plan {
case .superNode(let payload):
// super_node
await self.sendSuperPacket(type: .data, data: payload)
case .peer(let payload, let session):
// session
SDLLogger.log("[SDLContext] step 5 send packet by session: \(session)", for: .trace)
await self.sendPeerPacket(type: .data, data: payload, remoteAddress: session.natAddress)
self.flowTracer.inc(num: payload.count, type: .p2p)
case .superNodeAndPunch(let payload, let request):
// super_node
await self.sendSuperPacket(type: .data, data: payload)
SDLLogger.log("[SDLContext] step 5 send packet by super: \(self.config.stunSocketAddress)", for: .trace)
//
self.flowTracer.inc(num: payload.count, type: .forward)
//
await self.puncherActor.submitRegisterRequest(superClient: self.superClient, request: request)
}
}
}