调整Context的owns

This commit is contained in:
anlicheng 2026-05-07 09:41:26 +08:00
parent 5b37ac2552
commit 66e2877044
8 changed files with 708 additions and 572 deletions

View File

@ -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) {

View File

@ -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()
}

View File

@ -14,7 +14,7 @@ import NIOCore
1. rsa的加解密逻辑
*/
private func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task<Void, Never> {
func startMonitorTask(name: String, _ body: @escaping () async throws -> Void, retryDelay: Duration = .seconds(5)) -> Task<Void, Never> {
return Task(name: name) {
while true {
do {
@ -54,20 +54,6 @@ enum SDLContextError: Error {
actor SDLContextActor {
private enum UDPHoleKind: Equatable {
case v4
case v6
func convertAddressType() -> Session.AddressType {
switch self {
case .v4:
return .v4
case .v6:
return .v6
}
}
}
private var config: SDLConfiguration
// nat
var natType: SDLNATProberActor.NatType = .blocked
@ -82,26 +68,12 @@ actor SDLContextActor {
//
nonisolated let rsaCipher: RSACipher
//
private var udpHole: SDLUDPHole?
private var udpHoleMonitorTask: Task<Void, Never>?
private var natProbeTask: Task<Void, Never>?
private var udpHoleLocalAddress: SocketAddress?
private var udpHoleV6: SDLUDPHoleV6?
private var udpHoleV6MonitorTask: Task<Void, Never>?
private var udpHoleService: SDLUDPHoleService?
private var dnsService: SDLDNSService?
private var superService: SDLSuperService?
private var packetReaderService: SDLPacketReaderService?
// 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
//
@ -110,9 +82,6 @@ actor SDLContextActor {
// ipv6
private var ipv6AssistClient: SDLIPV6AssistClient?
//
private var readTask: Task<Void, Error>?
private let sessionManager = SessionManager()
nonisolated private let arpServer: ArpServer
@ -161,76 +130,33 @@ actor SDLContextActor {
await self.puncherActor.start()
await self.arpServer.start()
self.startDnsMonitor()
self.startDnsLocalMonitor()
self.startUDPHoleMonitor()
// self.startUDPHoleV6Monitor()
self.startSuperMonitor()
let dnsService = SDLDNSService(serverHost: self.config.serverHost, publicDnsServers: self.publicDnsServers) { [weak self] event in
await self?.handleDNSEvent(event)
}
self.dnsService = dnsService
await dnsService.start()
let udpHoleService = SDLUDPHoleService(proberActor: self.proberActor) { [weak self] event in
await self?.handleUDPHoleEvent(event)
}
self.udpHoleService = udpHoleService
await udpHoleService.start()
let superService = SDLSuperService(host: self.config.serverHost) { [weak self] message in
await self?.handleSuperMessage(message: message)
}
self.superService = superService
await superService.start()
}
// context
public func stop() async {
await self.puncherActor.stop()
await self.arpServer.clear()
await self.arpServer.stop()
await self.sessionManager.clear()
self.flowSessionManager.clear()
let udpHole = self.udpHole
self.udpHole = nil
self.udpHoleLocalAddress = nil
let udpHoleMonitorTask = self.udpHoleMonitorTask
self.udpHoleMonitorTask = nil
let natProbeTask = self.natProbeTask
self.natProbeTask = nil
udpHoleMonitorTask?.cancel()
natProbeTask?.cancel()
await self.proberActor.cancelAll()
await udpHole?.stop()
if let natProbeTask {
await natProbeTask.value
}
if let udpHoleMonitorTask {
await udpHoleMonitorTask.value
}
let udpHoleV6 = self.udpHoleV6
self.udpHoleV6 = nil
let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask
self.udpHoleV6MonitorTask = nil
udpHoleV6MonitorTask?.cancel()
udpHoleV6?.stop()
if let udpHoleV6MonitorTask {
await udpHoleV6MonitorTask.value
}
let dnsClient = self.dnsClient
self.dnsClient = nil
self.dnsMonitorTask?.cancel()
self.dnsMonitorTask = nil
dnsClient?.stop()
let dnsLocalClient = self.dnsLocalClient
self.dnsLocalClient = nil
self.dnsLocalMonitorTask?.cancel()
self.dnsLocalMonitorTask = nil
await dnsLocalClient?.stop()
let superClient = self.superClient
self.superClient = nil
self.superMonitorTask?.cancel()
self.superMonitorTask = nil
await superClient?.stop()
SDLLogger.log("[SDLContext] try to cancel readTask")
self.readTask?.cancel()
self.readTask = nil
self.registerTask?.cancel()
self.registerTask = nil
@ -240,6 +166,22 @@ actor SDLContextActor {
self.updatePolicyTask?.cancel()
self.updatePolicyTask = nil
let packetReaderService = self.packetReaderService
self.packetReaderService = nil
await packetReaderService?.stop()
let udpHoleService = self.udpHoleService
self.udpHoleService = nil
await udpHoleService?.stop()
let dnsService = self.dnsService
self.dnsService = nil
await dnsService?.stop()
let superService = self.superService
self.superService = nil
await superService?.stop()
self.sessionToken = nil
self.dataCipher = nil
self.natType = .blocked
@ -249,10 +191,6 @@ actor SDLContextActor {
}
deinit {
self.udpHole = nil
self.udpHoleLocalAddress = nil
self.udpHoleV6 = nil
self.dnsClient = nil
SDLLogger.log("[SDLContext] deinit", for: .debug)
}
@ -264,16 +202,6 @@ extension SDLContextActor {
private func setNatType(natType: SDLNATProberActor.NatType) {
self.natType = natType
}
//
private func probeNatType() async {
guard let udpHole = self.udpHole else {
return
}
// nat
self.natType = await self.proberActor.probeNatType(using: udpHole)
SDLLogger.log("[SDLContext] nat_type is: \(natType)")
}
}
// MARK: Notifier
@ -306,100 +234,13 @@ extension SDLContextActor {
}
private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
switch remoteAddress {
case .v4:
guard let udpHole = self.udpHole else {
SDLLogger.log("[SDLContext] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug)
return
}
await udpHole.send(type: type, data: data, remoteAddress: remoteAddress)
case .v6:
guard let udpHoleV6 = self.udpHoleV6 else {
SDLLogger.log("[SDLContext] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug)
return
}
udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress)
default:
SDLLogger.log("[SDLContext] unsupported socket family: \(remoteAddress)", for: .debug)
}
await self.udpHoleService?.send(type: type, data: data, remoteAddress: remoteAddress)
}
}
// MARK: Super
extension SDLContextActor {
private func startSuperMonitor() {
guard self.superMonitorTask == nil else {
return
}
self.superMonitorTask = startMonitorTask(name: "superMonitorTask") {
try await self.startSuperClient()
}
}
private func startSuperClient() async throws {
let superClient = SDLSuperClient(host: self.config.serverHost, port: 1443)
self.superClient = superClient
await superClient.start()
do {
try await withTaskCancellationHandler {
try await runSuperClient(superClient)
} onCancel: {
SDLLogger.log("[SDLContext] startSuperClient onCancel", for: .debug)
Task {
await superClient.stop()
}
}
await cleanupSuperClient(superClient)
} catch {
await cleanupSuperClient(superClient)
SDLLogger.log("[SDLContext] startSuperClient catch err: \(error)")
throw error
}
}
private func runSuperClient(_ superClient: SDLSuperClient) async throws {
try await Task.sleep(for: .seconds(0.5))
try Task.checkCancellation()
SDLLogger.log("[SDLContext] start super client: \(self.config.serverHost)")
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
for try await message in await superClient.messageStream {
try Task.checkCancellation()
await self.handleSuperMessage(message: message)
}
}
group.addTask {
while true {
try await Task.sleep(for: .seconds(5))
try Task.checkCancellation()
await superClient.send(type: .ping, data: Data())
}
}
_ = try await group.next()
}
}
private func cleanupSuperClient(_ superClient: SDLSuperClient) async {
await superClient.stop()
if self.superClient === superClient {
self.superClient = nil
}
SDLLogger.log("[SDLContext] cleanupSuperClient")
}
private func handleSuperMessage(message: SDLQUICInboundMessage) async {
switch message {
case .welcome(let welcome):
@ -443,7 +284,10 @@ extension SDLContextActor {
self.handleRegisterSuperNak(nakPacket: registerSuperNak)
case .peerInfo(let peerInfo):
SDLLogger.log("[SDLContext] peer message: \(peerInfo)")
await self.puncherActor.handlePeerInfo(using: self.udpHole, udpHoleV6: self.udpHoleV6, peerInfo: peerInfo)
let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo)
for packet in packets {
await self.udpHoleService?.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress)
}
case .event(let event):
await self.handleEvent(event: event)
case .policyReponse(let policyResponse):
@ -489,7 +333,7 @@ extension SDLContextActor {
do {
try await self.setNetworkSettings(config: self.config, dnsServer: DNSHelper.dnsServer)
SDLLogger.log("[SDLContext] setNetworkSettings successed")
self.startReader()
await self.startPacketReader()
//
await self.whenRegistedSuper()
} catch let err {
@ -507,7 +351,10 @@ extension SDLContextActor {
while true {
try await Task.sleep(for: .seconds(300))
SDLLogger.log("[SDLContext] updatePolicyTask execute")
await self.identifyStore.batUpdatePolicy(using: self.superClient, dstIdentityID: self.config.identityId)
let requests = await self.identifyStore.makeBatchPolicyRequests(dstIdentityID: self.config.identityId)
for request in requests {
await self.superService?.send(type: .policyRequest, data: request)
}
}
} catch let err {
SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)")
@ -573,171 +420,56 @@ extension SDLContextActor {
if let registerSuperData = try? registerSuper.serializedData() {
SDLLogger.log("[SDLContext] will send register super")
await self.superClient?.send(type: .registerSuper, data: registerSuperData)
await self.superService?.send(type: .registerSuper, data: registerSuperData)
}
}
}
// MARK: DnsLocal
// MARK: DNS service events
extension SDLContextActor {
private func startDnsLocalMonitor() {
guard self.dnsLocalMonitorTask == nil else {
return
}
self.dnsLocalMonitorTask = startMonitorTask(name: "dnsLocalMonitorTask") {
try await self.startDnsLocalClient()
}
}
private func startDnsLocalClient() async throws {
let dnsServer = self.publicDnsServers.randomElement() ?? self.publicDnsServers[0]
// dns
let dnsLocalClient = DNSLocalClient(host: dnsServer)
await dnsLocalClient.start()
SDLLogger.log("[SDLContext] dnsLocalClient started")
self.dnsLocalClient = dnsLocalClient
defer {
self.dnsLocalClient = nil
}
do {
try await withTaskCancellationHandler {
//
for try await packet in dnsLocalClient.packetFlow {
try Task.checkCancellation()
// Ip
let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
}
} onCancel: {
Task {
await dnsLocalClient.stop()
}
}
} catch let err {
await dnsLocalClient.stop()
throw err
}
}
}
// MARK: DnsCloud
extension SDLContextActor {
private func startDnsMonitor() {
guard self.dnsMonitorTask == nil else {
return
}
self.dnsMonitorTask = startMonitorTask(name: "dnsMonitorTask") {
try await self.startDnsClient()
}
}
private func startDnsClient() async throws {
// dns
let dnsClient = DNSCloudClient(host: self.config.serverHost, port: 15353)
self.dnsClient = dnsClient
dnsClient.start()
defer {
dnsClient.stop()
self.dnsClient = nil
}
try await withTaskCancellationHandler {
for try await packet in dnsClient.packetFlow {
try Task.checkCancellation()
let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
}
} onCancel: {
dnsClient.stop()
private func handleDNSEvent(_ event: SDLDNSService.Event) async {
switch event {
case .packet(let packet):
let nePacket = NEPacket(data: packet, protocolFamily: 2)
self.provider.packetFlow.writePacketObjects([nePacket])
}
}
}
// MARK: Hole
extension SDLContextActor {
private func startUDPHoleMonitor() {
guard self.udpHoleMonitorTask == nil else {
return
private func handleUDPHoleEvent(_ event: SDLUDPHoleService.Event) async {
switch event {
case .ready(let localAddress):
SDLLogger.log("[SDLContext] udpHole ready: \(localAddress)")
case .natType(let natType):
self.setNatType(natType: natType)
SDLLogger.log("[SDLContext] nat_type is: \(natType)")
case .packet(let remoteAddress, let message, let source):
await self.handleUDPHolePacket(remoteAddress: remoteAddress, message: message, source: source)
case .closed(let error):
SDLLogger.log("[SDLContext] udpHole closed: \(error)", for: .debug)
}
self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleMonitorTask") {
try await self.startUDPHole()
}
private func handleUDPHolePacket(remoteAddress: SocketAddress, message: SDLHoleMessage, source: SDLUDPHoleKind) async {
switch message.inboundMessage {
case .control(let message):
switch message {
case .stunReply(_), .stunProbeReply(_):
SDLLogger.log("[SDLContext] get a stun reply", for: .debug)
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: source)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: source)
}
case .data(let data):
try? await self.handleHoleData(data: data)
}
}
private func startUDPHole() async throws {
// udp
let udpHole = try SDLUDPHole()
let localAddress = try await udpHole.start()
SDLLogger.log("[SDLContext] udpHole started, on address: \(localAddress)")
self.udpHole = udpHole
self.udpHoleLocalAddress = localAddress
defer {
if self.udpHole === udpHole {
self.udpHole = nil
self.udpHoleLocalAddress = nil
}
}
// nat
self.natProbeTask?.cancel()
let proberActor = self.proberActor
self.natProbeTask = Task { [weak self] in
SDLLogger.log("[SDLContext] start probeNatType")
if Task.isCancelled {
return
}
let natType = await proberActor.probeNatType(using: udpHole)
if Task.isCancelled {
return
}
await self?.setNatType(natType: natType)
}
do {
try await withTaskCancellationHandler {
for try await (remoteAddress, message) in await udpHole.messageStream() {
try Task.checkCancellation()
switch message.inboundMessage {
case .control(let message):
switch message {
case .stunReply(_):
SDLLogger.log("[SDLContext] get a stunReply", for: .debug)
case .stunProbeReply(let probeReply):
await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply)
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v4)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
}
case .data(let data):
try? await self.handleHoleData(data: data)
}
}
} onCancel: {
Task {
await udpHole.stop()
}
}
} catch let err {
await udpHole.stop()
throw err
}
}
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: UDPHoleKind) async throws {
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: SDLUDPHoleKind) async throws {
let networkAddr = config.networkAddress
SDLLogger.log("[SDLContext] register packet: \(register), network_address: \(networkAddr)")
@ -761,7 +493,7 @@ extension SDLContextActor {
}
}
private func handleRegisterAck(remoteAddress: SocketAddress, registerAck: SDLRegisterAck, source: UDPHoleKind) async {
private func handleRegisterAck(remoteAddress: SocketAddress, registerAck: SDLRegisterAck, source: SDLUDPHoleKind) async {
// tun,
let networkAddr = config.networkAddress
if registerAck.dstMac == networkAddr.mac && registerAck.networkID == networkAddr.networkId {
@ -805,8 +537,9 @@ extension SDLContextActor {
SDLLogger.log("[SDLContext] hole identity: \(identityID), allow, data count: \(packetData.count)", for: .trace)
case .requestPolicy(let srcIdentityID):
SDLLogger.log("[SDLContext] not found identity: \(srcIdentityID) ruleMap", for: .debug)
//
await self.identifyStore.policyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.config.identityId, using: self.superClient)
if let queryData = await self.identifyStore.makePolicyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.config.identityId) {
await self.superService?.send(type: .policyRequest, data: queryData)
}
case .none:
()
}
@ -814,83 +547,6 @@ extension SDLContextActor {
}
// MARK: HoleV6
extension SDLContextActor {
private func startUDPHoleV6Monitor() {
guard self.udpHoleV6MonitorTask == nil else {
return
}
self.udpHoleV6MonitorTask = startMonitorTask(name: "udpHoleV6MonitorTask") {
try await self.startUDPHoleV6()
}
}
private func startUDPHoleV6() async throws {
// udp
let udpHoleV6 = try SDLUDPHoleV6()
let localAddress = try udpHoleV6.start()
self.udpHoleV6 = udpHoleV6
if let localAddress {
SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)")
} else {
SDLLogger.log("[SDLContext] udpHoleV6 started, no local address")
}
defer {
if self.udpHoleV6 === udpHoleV6 {
udpHoleV6.stop()
self.udpHoleV6 = nil
}
}
try await withThrowingTaskGroup { group in
defer {
group.cancelAll()
}
//
group.addTask {
for await (remoteAddress, message) in udpHoleV6.messageStream {
try Task.checkCancellation()
switch message.inboundMessage {
case .control(let message):
switch message {
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v6)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
default:
()
}
case .data(let data):
try? await self.handleHoleData(data: data)
}
}
}
group.addTask {
for await event in udpHoleV6.eventStream {
try Task.checkCancellation()
switch event {
case .ready:
SDLLogger.log("[SDLContext] udpHoleV6 ready")
case .closed, .errorCaught:
throw SDLContextError.udpHoleClosed
}
}
}
try await group.next()
}
}
}
// MARK: Stun
extension SDLContextActor {
@ -955,32 +611,20 @@ extension SDLContextActor {
extension SDLContextActor {
// , 线packetFlow
private func startReader() {
self.readTask?.cancel()
//
let provider = self.provider
self.readTask = Task(priority: .high) { [weak self] in
try await withTaskCancellationHandler {
do {
repeat {
try Task.checkCancellation()
let (packets, numbers) = await provider.packetFlow.readPackets()
try Task.checkCancellation()
for (data, number) in zip(packets, numbers) where number == 2 {
if let ipPacket = IPPacket(data) {
await self?.dealTunPacket(packet: ipPacket)
}
}
} while true
SDLLogger.log("[SDLContext] readTask finish")
} catch let err {
SDLLogger.log("[SDLContext] readTask catch error: \(err)")
throw err
}
} onCancel: {
SDLLogger.log("[SDLContext] readTask onCancel")
private func startPacketReader() async {
if self.packetReaderService == nil {
self.packetReaderService = SDLPacketReaderService(provider: self.provider) { [weak self] event in
await self?.handlePacketReaderEvent(event)
}
}
await self.packetReaderService?.start()
}
private func handlePacketReaderEvent(_ event: SDLPacketReaderService.Event) async {
switch event {
case .packet(let packet):
await self.dealTunPacket(packet: packet)
}
}
// ip: 0.0.0.0
@ -1080,10 +724,10 @@ extension SDLContextActor {
self.provider.packetFlow.writePacketObjects([nePacket])
case .cloudDNS(let name, let ipPacketData):
SDLLogger.log("[SDLContext] get cloud dns request: \(name)")
self.dnsClient?.forward(ipPacketData: ipPacketData)
await self.dnsService?.forward(ipPacketData: ipPacketData)
case .localDNS(let name, let payload, let tracker):
SDLLogger.log("[SDLContext] get local dns request: \(name)")
await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: payload)
await self.dnsService?.queryLocal(tracker: tracker, dnsPayload: payload)
case .forwardToNextHop(let ip, let type, let data, let kind):
await self.forwardPacketToNextHop(ip: ip, type: type, data: data, kind: kind)
case .drop(let reason):
@ -1106,11 +750,9 @@ extension SDLContextActor {
}
else {
SDLLogger.log("[SDLContext] dstIp: \(asIpAddress(ip)) arp query not found, broadcast", for: .trace)
// // arp广
// let arpReqeust = ARPPacket.arpRequest(senderIP: networkAddr.ip, senderMAC: networkAddr.mac, targetIP: dstIp)
// await self.routeLayerPacket(dstMac: ARPPacket.broadcastMac , type: .arp, data: arpReqeust.marshal())
try? await self.arpServer.arpRequest(targetIp: ip, use: self.superClient)
if let arpRequest = try? await self.arpServer.makeArpRequest(targetIp: ip) {
await self.superService?.send(type: .arpRequest, data: arpRequest)
}
}
}
@ -1149,7 +791,9 @@ extension SDLContextActor {
self.flowTracer.inc(num: payload.count, type: .forward)
//
await self.puncherActor.submitRegisterRequest(superClient: self.superClient, request: request)
if let queryData = await self.puncherActor.makeQueryInfoRequest(request: request) {
await self.superService?.send(type: .queryInfo, data: queryData)
}
}
}

View File

@ -0,0 +1,144 @@
import Foundation
import NetworkExtension
actor SDLDNSService {
enum Event {
case packet(Data)
}
typealias EventHandler = @Sendable (Event) async -> Void
private let serverHost: String
private let publicDnsServers: [String]
private let onEvent: EventHandler
private var dnsClient: DNSCloudClient?
private var dnsMonitorTask: Task<Void, Never>?
private var dnsLocalClient: DNSLocalClient?
private var dnsLocalMonitorTask: Task<Void, Never>?
init(serverHost: String, publicDnsServers: [String], onEvent: @escaping EventHandler) {
self.serverHost = serverHost
self.publicDnsServers = publicDnsServers
self.onEvent = onEvent
}
func start() {
self.startCloud()
self.startLocal()
}
func stop() async {
let dnsClient = self.dnsClient
self.dnsClient = nil
let dnsMonitorTask = self.dnsMonitorTask
self.dnsMonitorTask = nil
let dnsLocalClient = self.dnsLocalClient
self.dnsLocalClient = nil
let dnsLocalMonitorTask = self.dnsLocalMonitorTask
self.dnsLocalMonitorTask = nil
dnsMonitorTask?.cancel()
dnsClient?.stop()
if let dnsMonitorTask {
await dnsMonitorTask.value
}
dnsLocalMonitorTask?.cancel()
await dnsLocalClient?.stop()
if let dnsLocalMonitorTask {
await dnsLocalMonitorTask.value
}
}
func forward(ipPacketData: Data) {
self.dnsClient?.forward(ipPacketData: ipPacketData)
}
func queryLocal(tracker: DNSLocalClient.DNSTracker, dnsPayload: Data) async {
await self.dnsLocalClient?.query(tracker: tracker, dnsPayload: dnsPayload)
}
private func startCloud() {
guard self.dnsMonitorTask == nil else {
return
}
self.dnsMonitorTask = startMonitorTask(name: "dnsServiceCloudMonitor") { [weak self] in
guard let self else {
throw CancellationError()
}
try await self.runCloud()
}
}
private func startLocal() {
guard self.dnsLocalMonitorTask == nil else {
return
}
self.dnsLocalMonitorTask = startMonitorTask(name: "dnsServiceLocalMonitor") { [weak self] in
guard let self else {
throw CancellationError()
}
try await self.runLocal()
}
}
private func runCloud() async throws {
let dnsClient = DNSCloudClient(host: self.serverHost, port: 15353)
self.dnsClient = dnsClient
dnsClient.start()
defer {
dnsClient.stop()
if self.dnsClient === dnsClient {
self.dnsClient = nil
}
}
let onEvent = self.onEvent
try await withTaskCancellationHandler {
for try await packet in dnsClient.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
} onCancel: {
dnsClient.stop()
}
}
private func runLocal() async throws {
let dnsServer = self.publicDnsServers.randomElement() ?? "223.5.5.5"
let dnsLocalClient = DNSLocalClient(host: dnsServer)
await dnsLocalClient.start()
self.dnsLocalClient = dnsLocalClient
SDLLogger.log("[SDLDNSService] dnsLocalClient started")
defer {
if self.dnsLocalClient === dnsLocalClient {
self.dnsLocalClient = nil
}
}
let onEvent = self.onEvent
do {
try await withTaskCancellationHandler {
for try await packet in dnsLocalClient.packetFlow {
try Task.checkCancellation()
await onEvent(.packet(packet))
}
} onCancel: {
Task {
await dnsLocalClient.stop()
}
}
await dnsLocalClient.stop()
} catch {
await dnsLocalClient.stop()
throw error
}
}
}

View File

@ -0,0 +1,49 @@
import Foundation
import NetworkExtension
actor SDLPacketReaderService {
enum Event {
case packet(IPPacket)
}
typealias EventHandler = @Sendable (Event) async -> Void
private let provider: NEPacketTunnelProvider
private let onEvent: EventHandler
private var readTask: Task<Void, Never>?
init(provider: NEPacketTunnelProvider, onEvent: @escaping EventHandler) {
self.provider = provider
self.onEvent = onEvent
}
func start() {
guard self.readTask == nil else {
return
}
let provider = self.provider
let onEvent = self.onEvent
self.readTask = Task(priority: .high) {
while !Task.isCancelled {
let (packets, numbers) = await provider.packetFlow.readPackets()
if Task.isCancelled {
break
}
for (data, number) in zip(packets, numbers) where number == 2 {
if let packet = IPPacket(data) {
await onEvent(.packet(packet))
}
}
}
SDLLogger.log("[SDLPacketReaderService] readTask finished")
}
}
func stop() {
let readTask = self.readTask
self.readTask = nil
readTask?.cancel()
}
}

View File

@ -0,0 +1,111 @@
import Foundation
actor SDLSuperService {
typealias MessageHandler = @Sendable (SDLQUICInboundMessage) async -> Void
private let host: String
private let port: UInt16
private let onMessage: MessageHandler
private var superClient: SDLSuperClient?
private var monitorTask: Task<Void, Never>?
init(host: String, port: UInt16 = 1443, onMessage: @escaping MessageHandler) {
self.host = host
self.port = port
self.onMessage = onMessage
}
func start() {
guard self.monitorTask == nil else {
return
}
self.monitorTask = startMonitorTask(name: "superServiceMonitor") { [weak self] in
guard let self else {
throw CancellationError()
}
try await self.runOnce()
}
}
func stop() async {
let monitorTask = self.monitorTask
self.monitorTask = nil
let superClient = self.superClient
self.superClient = nil
monitorTask?.cancel()
await superClient?.stop()
if let monitorTask {
await monitorTask.value
}
}
func send(type: SDLPacketType, data: Data) async {
await self.superClient?.send(type: type, data: data)
}
private func runOnce() async throws {
let superClient = SDLSuperClient(host: self.host, port: self.port)
self.superClient = superClient
await superClient.start()
do {
try await withTaskCancellationHandler {
try await self.run(superClient)
} onCancel: {
Task {
await superClient.stop()
}
}
await self.cleanup(superClient)
} catch {
await self.cleanup(superClient)
throw error
}
}
private func run(_ superClient: SDLSuperClient) async throws {
try await Task.sleep(for: .seconds(0.5))
try Task.checkCancellation()
SDLLogger.log("[SDLSuperService] start super client: \(self.host)")
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
let onMessage = self.onMessage
group.addTask {
for try await message in await superClient.messageStream {
try Task.checkCancellation()
await onMessage(message)
}
}
group.addTask {
while true {
try await Task.sleep(for: .seconds(5))
try Task.checkCancellation()
await superClient.send(type: .ping, data: Data())
}
}
_ = try await group.next()
}
}
private func cleanup(_ superClient: SDLSuperClient) async {
await superClient.stop()
if self.superClient === superClient {
self.superClient = nil
}
SDLLogger.log("[SDLSuperService] cleanup")
}
}

View File

@ -0,0 +1,240 @@
import Foundation
import NIOCore
enum SDLUDPHoleKind: Equatable {
case v4
case v6
func convertAddressType() -> Session.AddressType {
switch self {
case .v4:
return .v4
case .v6:
return .v6
}
}
}
actor SDLUDPHoleService {
enum Event {
case ready(SocketAddress)
case natType(SDLNATProberActor.NatType)
case packet(SocketAddress, SDLHoleMessage, source: SDLUDPHoleKind)
case closed(Error)
}
typealias EventHandler = @Sendable (Event) async -> Void
private let proberActor: SDLNATProberActor
private let onEvent: EventHandler
private var udpHole: SDLUDPHole?
private var udpHoleMonitorTask: Task<Void, Never>?
private var natProbeTask: Task<Void, Never>?
private var localAddress: SocketAddress?
private var udpHoleV6: SDLUDPHoleV6?
private var udpHoleV6MonitorTask: Task<Void, Never>?
init(proberActor: SDLNATProberActor, onEvent: @escaping EventHandler) {
self.proberActor = proberActor
self.onEvent = onEvent
}
func start(includeV6: Bool = false) {
self.startV4()
if includeV6 {
self.startV6()
}
}
func stop() async {
let udpHole = self.udpHole
self.udpHole = nil
self.localAddress = nil
let udpHoleMonitorTask = self.udpHoleMonitorTask
self.udpHoleMonitorTask = nil
let natProbeTask = self.natProbeTask
self.natProbeTask = nil
udpHoleMonitorTask?.cancel()
natProbeTask?.cancel()
await self.proberActor.cancelAll()
await udpHole?.stop()
if let natProbeTask {
await natProbeTask.value
}
if let udpHoleMonitorTask {
await udpHoleMonitorTask.value
}
let udpHoleV6 = self.udpHoleV6
self.udpHoleV6 = nil
let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask
self.udpHoleV6MonitorTask = nil
udpHoleV6MonitorTask?.cancel()
udpHoleV6?.stop()
if let udpHoleV6MonitorTask {
await udpHoleV6MonitorTask.value
}
}
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
switch remoteAddress {
case .v4:
guard let udpHole else {
SDLLogger.log("[SDLUDPHoleService] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug)
return
}
await udpHole.send(type: type, data: data, remoteAddress: remoteAddress)
case .v6:
guard let udpHoleV6 else {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug)
return
}
udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress)
default:
SDLLogger.log("[SDLUDPHoleService] unsupported socket family: \(remoteAddress)", for: .debug)
}
}
private func startV4() {
guard self.udpHoleMonitorTask == nil else {
return
}
self.udpHoleMonitorTask = startMonitorTask(name: "udpHoleServiceV4Monitor") { [weak self] in
guard let self else {
throw CancellationError()
}
try await self.runV4()
}
}
private func runV4() async throws {
let udpHole = try SDLUDPHole()
let localAddress = try await udpHole.start()
self.udpHole = udpHole
self.localAddress = localAddress
SDLLogger.log("[SDLUDPHoleService] udpHole started, on address: \(localAddress)")
await self.onEvent(.ready(localAddress))
self.startNatProbe(using: udpHole)
defer {
if self.udpHole === udpHole {
self.udpHole = nil
self.localAddress = nil
}
}
do {
try await withTaskCancellationHandler {
for try await (remoteAddress, message) in await udpHole.messageStream() {
try Task.checkCancellation()
try await self.handleV4Message(remoteAddress: remoteAddress, message: message)
}
} onCancel: {
Task {
await udpHole.stop()
}
}
} catch {
await udpHole.stop()
throw error
}
}
private func startNatProbe(using udpHole: SDLUDPHole) {
self.natProbeTask?.cancel()
let proberActor = self.proberActor
let onEvent = self.onEvent
self.natProbeTask = Task {
if Task.isCancelled {
return
}
let natType = await proberActor.probeNatType(using: udpHole)
if Task.isCancelled {
return
}
await onEvent(.natType(natType))
}
}
private func handleV4Message(remoteAddress: SocketAddress, message: SDLHoleMessage) async throws {
switch message.inboundMessage {
case .control(let control):
switch control {
case .stunProbeReply(let probeReply):
await self.proberActor.handleProbeReply(localAddress: self.localAddress, reply: probeReply)
default:
await self.onEvent(.packet(remoteAddress, message, source: .v4))
}
case .data:
await self.onEvent(.packet(remoteAddress, message, source: .v4))
}
}
private func startV6() {
guard self.udpHoleV6MonitorTask == nil else {
return
}
self.udpHoleV6MonitorTask = startMonitorTask(name: "udpHoleServiceV6Monitor") { [weak self] in
guard let self else {
throw CancellationError()
}
try await self.runV6()
}
}
private func runV6() async throws {
let udpHoleV6 = try SDLUDPHoleV6()
let localAddress = try udpHoleV6.start()
self.udpHoleV6 = udpHoleV6
if let localAddress {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, on address: \(localAddress)")
} else {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, no local address")
}
defer {
if self.udpHoleV6 === udpHoleV6 {
udpHoleV6.stop()
self.udpHoleV6 = nil
}
}
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
let onEvent = self.onEvent
group.addTask {
for await (remoteAddress, message) in udpHoleV6.messageStream {
try Task.checkCancellation()
await onEvent(.packet(remoteAddress, message, source: .v6))
}
}
group.addTask {
for await event in udpHoleV6.eventStream {
try Task.checkCancellation()
switch event {
case .ready:
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 ready")
case .closed, .errorCaught:
throw SDLContextError.udpHoleClosed
}
}
}
_ = try await group.next()
}
}
}

View File

@ -24,49 +24,35 @@ actor IdentityStore {
init(publisher: SnapshotPublisher<IdentitySnapshot>) {
self.publisher = publisher
}
// , superClient
func batUpdatePolicy(using superClient: SDLSuperClient?, dstIdentityID: UInt32) async {
guard let superClient else {
return
}
for identityId in self.identityMap.keys {
func makeBatchPolicyRequests(dstIdentityID: UInt32) -> [Data] {
return self.identityMap.keys.compactMap { identityId in
var policyRequest = SDLPolicyRequest()
policyRequest.srcIdentityID = identityId
policyRequest.dstIdentityID = dstIdentityID
policyRequest.version = self.nextVersion(identityId: identityId)
//
if let queryData = try? policyRequest.serializedData() {
await superClient.send(type: .policyRequest, data: queryData)
}
return try? policyRequest.serializedData()
}
}
//
func policyRequest(srcIdentityId: UInt32, dstIdentityId: UInt32, using superClient: SDLSuperClient?) async {
guard let superClient, !coolingDown.contains(srcIdentityId) else {
return
func makePolicyRequest(srcIdentityId: UInt32, dstIdentityId: UInt32) -> Data? {
guard !coolingDown.contains(srcIdentityId) else {
return nil
}
var policyRequest = SDLPolicyRequest()
policyRequest.srcIdentityID = srcIdentityId
policyRequest.dstIdentityID = dstIdentityId
policyRequest.version = self.nextVersion(identityId: srcIdentityId)
//
coolingDown.insert(srcIdentityId)
//
if let queryData = try? policyRequest.serializedData() {
await superClient.send(type: .policyRequest, data: queryData)
}
Task {
//
Task { [weak self] in
try? await Task.sleep(for: .seconds(5))
self.endCooldown(for: srcIdentityId)
await self?.endCooldown(for: srcIdentityId)
}
return try? policyRequest.serializedData()
}
//