punchnet-macos/Tun/Context/SDLContextActor.swift
2026-05-27 21:24:22 +08:00

759 lines
28 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
*/
enum SDLContextError: Error {
case udpHoleClosed
case dnsLocalClientClosed
case dnsLocalClientCancelled
case dnsClientClosed
case dnsClientCancelled
}
actor SDLContextActor {
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 dnsService: DNSService?
private let superService: SDLSuperService
private let udpHoleService: SDLUDPHoleService
private let udpHoleV6Service: SDLUDPHoleV6Service
private let packetOutboundActor: PacketOutboundActor
private let packetInboundActor: PacketInboundActor
private let tunNetworkManager: SDLTunNetworkManager
private let publicDnsServers = ["223.5.5.5", "119.29.29.29"]
nonisolated private let puncherActor: SDLPuncherActor
//
nonisolated private let proberActor: SDLNATProberActor
// ipv6
private var ipv6AssistClient: SDLIPV6AssistClient?
private let ipv6AssistEvents: AsyncStream<SDLV6Info?>
private let ipv6AssistContinuation: AsyncStream<SDLV6Info?>.Continuation
private let sessionManager: SessionManager
nonisolated private let arpResolver: ArpResolver
// socket
// App Group + Darwin Notification
//
nonisolated private let flowTracer: SDLFlowTracer
nonisolated private let provider: NEPacketTunnelProvider
//
private let policyService: PolicyService
private var rootTask: Task<Void, Error>?
private let readySignal = AsyncOneShot<Void>()
public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) {
let puncherActor = SDLPuncherActor()
let proberActor = SDLNATProberActor(addressArray: config.stunProbeSocketAddressArray)
let sessionManager = SessionManager()
let arpResolver = ArpResolver()
let flowTracer = SDLFlowTracer()
let policyService = PolicyService(identityId: config.identityId, acl: config.acl)
let superService = SDLSuperService(serverEndpoint: config.serverEndpoint)
let udpHoleService = SDLUDPHoleService(proberActor: proberActor)
let udpHoleV6Service = SDLUDPHoleV6Service()
let tunNetworkManager = SDLTunNetworkManager(provider: provider)
let ipv6AssistPair = AsyncStream.makeStream(of: Optional<SDLV6Info>.self, bufferingPolicy: .bufferingNewest(1))
let packetOutboundActor = PacketOutboundActor(
provider: provider,
config: config,
dataCipher: nil,
sessionManager: sessionManager,
arpResolver: arpResolver,
puncherActor: puncherActor,
policyService: policyService,
superService: superService,
udpHoleService: udpHoleService,
udpHoleV6Service: udpHoleV6Service,
flowTracer: flowTracer
)
let packetInboundActor = PacketInboundActor(
provider: provider,
config: config,
dataCipher: nil,
policyService: policyService,
packetOutboundActor: packetOutboundActor,
arpResolver: arpResolver,
superService: superService,
flowTracer: flowTracer
)
self.provider = provider
self.config = config
self.rsaCipher = rsaCipher
self.puncherActor = puncherActor
self.proberActor = proberActor
self.sessionManager = sessionManager
self.arpResolver = arpResolver
self.flowTracer = flowTracer
//
self.policyService = policyService
self.superService = superService
self.udpHoleService = udpHoleService
self.udpHoleV6Service = udpHoleV6Service
self.packetOutboundActor = packetOutboundActor
self.packetInboundActor = packetInboundActor
self.tunNetworkManager = tunNetworkManager
self.ipv6AssistEvents = ipv6AssistPair.stream
self.ipv6AssistContinuation = ipv6AssistPair.continuation
}
public func start() async throws {
guard self.rootTask == nil else {
try await self.readySignal.wait(timeout: .seconds(30))
return
}
let rootTask = Task {
try await self.runRoot()
}
self.rootTask = rootTask
do {
try await self.readySignal.wait(timeout: .seconds(30))
} catch {
rootTask.cancel()
_ = try? await rootTask.value
self.rootTask = nil
throw error
}
}
// context
public func stop() async {
let rootTask = self.rootTask
self.rootTask = nil
rootTask?.cancel()
_ = try? await rootTask?.value
await self.cleanupRoot()
}
private func runRoot() async throws {
do {
try await self.runRootBody()
await self.cleanupRoot()
} catch is CancellationError {
await self.cleanupRoot()
throw CancellationError()
} catch {
await self.readySignal.fail(error)
await self.cleanupRoot()
throw error
}
}
private func runRootBody() async throws {
self.prepareTunnelNotifier()
let dnsService = DNSService(serverIP: self.config.serverEndpoint.ip, publicDnsServers: self.publicDnsServers) { [weak self] event in
await self?.handleDNSEvent(event)
}
self.dnsService = dnsService
await self.packetOutboundActor.updateDNSService(dnsService)
await self.superService.updateMessageHandler { [weak self] message in
await self?.handleSuperMessage(message: message)
}
let packetInboundActor = self.packetInboundActor
await self.udpHoleService.updateHandlers(
onEvent: { [weak self] event in
await self?.handleUDPHoleControlEvent(event)
},
onData: { data in
await packetInboundActor.handleData(data)
}
)
await self.udpHoleV6Service.updateHandlers(
onEvent: { [weak self] event in
await self?.handleUDPHoleControlEvent(event)
},
onData: { data in
await packetInboundActor.handleData(data)
}
)
let superService = self.superService
let udpHoleService = self.udpHoleService
let udpHoleV6Service = self.udpHoleV6Service
let packetOutboundActor = self.packetOutboundActor
let policyService = self.policyService
let puncherActor = self.puncherActor
let arpResolver = self.arpResolver
let readySignal = self.readySignal
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await Self.runRestarting(name: "superService") {
try await superService.run()
}
}
group.addTask {
try await Self.runRestarting(name: "udpHoleService") {
try await udpHoleService.run()
}
}
group.addTask {
try await Self.runRestarting(name: "udpHoleV6Service") {
try await udpHoleV6Service.run()
}
}
group.addTask {
try await dnsService.run()
}
group.addTask {
try await puncherActor.runCleanup()
}
group.addTask {
try await arpResolver.runCleanup()
}
group.addTask {
try await self.runIPv6AssistSupervisor()
}
group.addTask(priority: .high) {
_ = try await readySignal.wait()
try await packetOutboundActor.runPacketReader()
}
group.addTask {
_ = try await readySignal.wait()
try await Self.runPeriodic(name: "updatePolicyTask", interval: .seconds(10)) {
SDLLogger.log("[SDLContext] updatePolicyTask execute")
await policyService.updatePolicy(superService: superService)
}
}
group.addTask {
_ = try await readySignal.wait()
try await Self.runPeriodic(name: "stunRequestTask", interval: .seconds(8)) {
try await self.runStunRequestOnce()
}
}
try await group.waitForAll()
}
}
private static func runRestarting(
name: String,
retryDelay: Duration = .seconds(5),
operation: @escaping @Sendable () async throws -> Void
) async throws {
while !Task.isCancelled {
do {
try Task.checkCancellation()
try await operation()
SDLLogger.log("[SDLContext] worker \(name) ended, will restart", for: .debug)
} catch is CancellationError {
SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug)
throw CancellationError()
} catch {
SDLLogger.log("[SDLContext] worker \(name) crashed: \(error.localizedDescription), will restart", for: .debug)
}
try await Task.sleep(for: retryDelay)
}
}
private static func runPeriodic(
name: String,
interval: Duration,
retryDelay: Duration = .seconds(5),
operation: @escaping @Sendable () async throws -> Void
) async throws {
while !Task.isCancelled {
do {
try Task.checkCancellation()
try await operation()
try await Task.sleep(for: interval)
} catch is CancellationError {
SDLLogger.log("[SDLContext] worker \(name) cancelled", for: .debug)
throw CancellationError()
} catch {
SDLLogger.log("[SDLContext] worker \(name) crashed: \(error.localizedDescription), will retry", for: .debug)
try await Task.sleep(for: retryDelay)
}
}
}
private func runIPv6AssistSupervisor() async throws {
do {
for await assistInfo in self.ipv6AssistEvents {
try Task.checkCancellation()
await self.stopCurrentIPv6AssistClient()
guard let assistInfo else {
continue
}
guard let client = SDLIPV6AssistClient(assistServerInfo: assistInfo) else {
SDLLogger.log("[SDLContext] invalid ipv6 assist config", for: .debug)
continue
}
self.ipv6AssistClient = client
do {
try await client.run()
} catch is CancellationError {
throw CancellationError()
} catch {
SDLLogger.log("[SDLContext] ipv6 assist client ended: \(error.localizedDescription)", for: .debug)
}
if self.ipv6AssistClient === client {
self.ipv6AssistClient = nil
}
}
} catch is CancellationError {
await self.stopCurrentIPv6AssistClient()
throw CancellationError()
}
await self.stopCurrentIPv6AssistClient()
}
private func stopCurrentIPv6AssistClient() async {
let client = self.ipv6AssistClient
self.ipv6AssistClient = nil
await client?.stop()
}
private func cleanupRoot() async {
await self.puncherActor.stop()
await self.arpResolver.stop()
await self.sessionManager.clear()
await self.policyService.clear()
await self.udpHoleService.stop()
await self.udpHoleV6Service.stop()
let dnsService = self.dnsService
self.dnsService = nil
await self.packetOutboundActor.updateDNSService(nil)
await dnsService?.stop()
await self.superService.stop()
self.sessionToken = nil
self.dataCipher = nil
self.natType = .blocked
await self.packetOutboundActor.updateRuntime(config: self.config, dataCipher: nil)
await self.packetInboundActor.updateRuntime(config: self.config, dataCipher: nil)
self.ipv6AssistContinuation.yield(nil)
await self.stopCurrentIPv6AssistClient()
}
deinit {
SDLLogger.log("[SDLContext] deinit", for: .debug)
}
}
// MARK: probe
extension SDLContextActor {
private func setNatType(natType: SDLNATProberActor.NatType) {
self.natType = 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:
await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress)
case .v6:
await self.udpHoleV6Service.send(type: type, data: data, remoteAddress: remoteAddress)
default:
SDLLogger.log("[SDLContext] unsupported socket family: \(remoteAddress)", for: .debug)
}
}
}
// MARK: Super
extension SDLContextActor {
private func handleSuperMessage(message: SDLQUICInboundMessage) async {
switch message {
case .welcome(let welcome):
SDLLogger.log("[SDLContext] quic welcome: \(welcome)")
await self.stopCurrentIPv6AssistClient()
// v6
if welcome.hasIpv6Assist {
self.ipv6AssistContinuation.yield(welcome.ipv6Assist)
} else {
self.ipv6AssistContinuation.yield(nil)
}
//
await self.doRegisterSuper()
SDLLogger.log("[SDLContext] quic doRegisterSuper")
case .pong:
//SDLLogger.shared.log("[SDLContext] quic pong")
()
case .registerSuperAck(let registerSuperAck):
await self.handleRegisterSuperAck(registerSuperAck: registerSuperAck)
case .registerSuperNak(let registerSuperNak):
await self.handleRegisterSuperNak(nakPacket: registerSuperNak)
case .peerInfo(let peerInfo):
SDLLogger.log("[SDLContext] peer message: \(peerInfo)")
let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo)
for packet in packets {
await self.sendPacket(type: .register, data: packet.data, remoteAddress: packet.remoteAddress)
}
case .event(let event):
await self.handleEvent(event: event)
case .policyReponse(let policyResponse):
//
await self.policyService.applyPolicyResponse(policyResponse)
case .exposedServiceResponse(let response):
await self.applyExposedServiceResponse(response)
case .arpResponse(let arpResponse):
SDLLogger.log("[SDLContext] get arp response: \(arpResponse)")
await self.arpResolver.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
await self.readySignal.fail(error)
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)
await self.readySignal.fail(error)
self.provider.cancelTunnelWithError(error)
return
}
await self.packetOutboundActor.updateRuntime(config: self.config, dataCipher: self.dataCipher)
await self.packetInboundActor.updateRuntime(config: self.config, dataCipher: self.dataCipher)
SDLLogger.log("[SDLContext] registerSuperAck, use algorithm \(algorithm), key len: \(key.count)")
// tun
do {
try await self.tunNetworkManager.apply(settings: .init(config: self.config), dnsServer: DNSHelper.dnsServer)
SDLLogger.log("[SDLContext] setNetworkSettings successed")
await self.readySignal.succeed(())
} catch let err {
SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)")
await self.readySignal.fail(err)
self.provider.cancelTunnelWithError(err)
}
}
private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) async {
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)
await self.readySignal.fail(error)
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 .requestExposedService:
await self.requestExposedService()
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.superService.send(type: .registerSuper, data: registerSuperData)
}
}
private func requestExposedService() async {
guard let requestData = await self.policyService.makeExposedServiceRequest() else {
return
}
await self.superService.send(type: .exposedServiceRequest, data: requestData)
}
private func applyExposedServiceResponse(_ response: SDLExposedServiceResponse) async {
guard let acl = await self.policyService.applyExposedServiceResponse(response) else {
return
}
self.config.acl = acl
}
}
// MARK: DNS service events
extension SDLContextActor {
private func handleDNSEvent(_ event: DNSService.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 handleUDPHoleControlEvent(_ 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)
}
}
private func handleUDPHolePacket(remoteAddress: SocketAddress, message: SDLHoleControlMessage, source: SDLUDPHoleKind) async {
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)
}
}
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: SDLUDPHoleKind) 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: SDLUDPHoleKind) 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)")
}
}
}
// MARK: Stun
extension SDLContextActor {
private func runStunRequestOnce() async throws {
let probeReply = try? await self.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(v6Info: probeReply?.v6Info)
}
private func sendStunRequest(v6Info: SDLV6Info?) 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 {
stunRequest.v6Info = v6Info
}
if let stunData = try? stunRequest.serializedData() {
await self.sendSuperPacket(type: .stunRequest, data: stunData)
}
}
}
// MARK: NEPacketTunnelProvider
extension SDLContextActor {
// 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
}
await self.packetOutboundActor.updateRuntime(config: self.config, dataCipher: self.dataCipher)
await self.packetInboundActor.updateRuntime(config: self.config, dataCipher: self.dataCipher)
try await self.tunNetworkManager.apply(settings: .init(config: self.config), dnsServer: DNSHelper.dnsServer)
}
}