修复流程

This commit is contained in:
anlicheng 2026-05-27 15:17:05 +08:00
parent 80d5c7dfb1
commit 8183e3b3bf
7 changed files with 463 additions and 252 deletions

View File

@ -83,7 +83,7 @@ final class SDLRuntimeEnvironment {
) )
self.contextActor = contextActor self.contextActor = contextActor
await contextActor.start() try await contextActor.start()
self.state = .running self.state = .running
case .running: case .running:
SDLLogger.log("[SDLRuntimeEnvironment] is running, ignore start command") SDLLogger.log("[SDLRuntimeEnvironment] is running, ignore start command")

View File

@ -64,8 +64,8 @@ actor SDLContextActor {
nonisolated let rsaCipher: RSACipher nonisolated let rsaCipher: RSACipher
private var dnsService: DNSService? private var dnsService: DNSService?
private let superServiceProxy: SDLSuperServiceProxy private let superService: SDLSuperService
private let udpHoleServiceProxy: SDLUDPHoleServiceProxy private let udpHoleService: SDLUDPHoleService
private let packetOutboundActor: PacketOutboundActor private let packetOutboundActor: PacketOutboundActor
private let packetInboundActor: PacketInboundActor private let packetInboundActor: PacketInboundActor
private let tunNetworkManager: SDLTunNetworkManager private let tunNetworkManager: SDLTunNetworkManager
@ -96,6 +96,8 @@ actor SDLContextActor {
// stunRequest // stunRequest
private var stunRequestWorker: PeriodicWorker? private var stunRequestWorker: PeriodicWorker?
private var rootTask: Task<Void, Error>?
private let readySignal = AsyncOneShot<Void>()
public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) { public init(provider: NEPacketTunnelProvider, config: SDLConfiguration, rsaCipher: RSACipher) {
let puncherActor = SDLPuncherActor() let puncherActor = SDLPuncherActor()
@ -104,8 +106,8 @@ actor SDLContextActor {
let arpResolver = ArpResolver() let arpResolver = ArpResolver()
let flowTracer = SDLFlowTracer() let flowTracer = SDLFlowTracer()
let policyService = PolicyService(identityId: config.identityId, acl: config.acl) let policyService = PolicyService(identityId: config.identityId, acl: config.acl)
let superServiceProxy = SDLSuperServiceProxy() let superService = SDLSuperService(serverEndpoint: config.serverEndpoint)
let udpHoleServiceProxy = SDLUDPHoleServiceProxy() let udpHoleService = SDLUDPHoleService(proberActor: proberActor)
let tunNetworkManager = SDLTunNetworkManager(provider: provider) let tunNetworkManager = SDLTunNetworkManager(provider: provider)
let packetOutboundActor = PacketOutboundActor( let packetOutboundActor = PacketOutboundActor(
provider: provider, provider: provider,
@ -115,8 +117,8 @@ actor SDLContextActor {
arpResolver: arpResolver, arpResolver: arpResolver,
puncherActor: puncherActor, puncherActor: puncherActor,
policyService: policyService, policyService: policyService,
superServiceProxy: superServiceProxy, superService: superService,
udpHoleServiceProxy: udpHoleServiceProxy, udpHoleService: udpHoleService,
flowTracer: flowTracer flowTracer: flowTracer
) )
let packetInboundActor = PacketInboundActor( let packetInboundActor = PacketInboundActor(
@ -126,7 +128,7 @@ actor SDLContextActor {
policyService: policyService, policyService: policyService,
packetOutboundActor: packetOutboundActor, packetOutboundActor: packetOutboundActor,
arpResolver: arpResolver, arpResolver: arpResolver,
superServiceProxy: superServiceProxy, superService: superService,
flowTracer: flowTracer flowTracer: flowTracer
) )
@ -143,14 +145,60 @@ actor SDLContextActor {
// //
self.policyService = policyService self.policyService = policyService
self.superServiceProxy = superServiceProxy self.superService = superService
self.udpHoleServiceProxy = udpHoleServiceProxy self.udpHoleService = udpHoleService
self.packetOutboundActor = packetOutboundActor self.packetOutboundActor = packetOutboundActor
self.packetInboundActor = packetInboundActor self.packetInboundActor = packetInboundActor
self.tunNetworkManager = tunNetworkManager self.tunNetworkManager = tunNetworkManager
} }
public func start() async { 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() self.prepareTunnelNotifier()
// arp // arp
@ -164,29 +212,66 @@ actor SDLContextActor {
await self.packetOutboundActor.updateDNSService(dnsService) await self.packetOutboundActor.updateDNSService(dnsService)
await dnsService.start() await dnsService.start()
let udpHoleEventHandler = await self.udpHoleServiceProxy.makeEventHandler { [weak self] event in await self.superService.updateMessageHandler { [weak self] message in
await self?.handleUDPHoleControlEvent(event) await self?.handleSuperMessage(message: message)
} }
let packetInboundActor = self.packetInboundActor let packetInboundActor = self.packetInboundActor
let udpHoleService = SDLUDPHoleService( await self.udpHoleService.updateHandlers(
proberActor: self.proberActor, onEvent: { [weak self] event in
onEvent: udpHoleEventHandler, await self?.handleUDPHoleControlEvent(event)
},
onData: { data in onData: { data in
await packetInboundActor.handleData(data) await packetInboundActor.handleData(data)
} }
) )
await self.udpHoleServiceProxy.replace(udpHoleService)
await udpHoleService.start(includeV6: false)
let superService = SDLSuperService(serverEndpoint: self.config.serverEndpoint) { [weak self] message in let superService = self.superService
await self?.handleSuperMessage(message: message) let udpHoleService = self.udpHoleService
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(includeV6: false)
}
}
try await group.waitForAll()
} }
await self.superServiceProxy.replace(superService)
await superService.start()
} }
// context private static func runRestarting(
public func stop() async { 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 func cleanupRoot() async {
await self.puncherActor.stop() await self.puncherActor.stop()
await self.arpResolver.stop() await self.arpResolver.stop()
await self.sessionManager.clear() await self.sessionManager.clear()
@ -199,15 +284,14 @@ actor SDLContextActor {
await self.policyService.clear() await self.policyService.clear()
await self.packetOutboundActor.stop() await self.packetOutboundActor.stop()
await self.udpHoleService.stop()
await self.udpHoleServiceProxy.stop()
let dnsService = self.dnsService let dnsService = self.dnsService
self.dnsService = nil self.dnsService = nil
await self.packetOutboundActor.updateDNSService(nil) await self.packetOutboundActor.updateDNSService(nil)
await dnsService?.stop() await dnsService?.stop()
await self.superServiceProxy.stop() await self.superService.stop()
self.sessionToken = nil self.sessionToken = nil
self.dataCipher = nil self.dataCipher = nil
@ -263,7 +347,7 @@ extension SDLContextActor {
} }
private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
await self.udpHoleServiceProxy.send(type: type, data: data, remoteAddress: remoteAddress) await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress)
} }
} }
@ -292,12 +376,12 @@ extension SDLContextActor {
case .registerSuperAck(let registerSuperAck): case .registerSuperAck(let registerSuperAck):
await self.handleRegisterSuperAck(registerSuperAck: registerSuperAck) await self.handleRegisterSuperAck(registerSuperAck: registerSuperAck)
case .registerSuperNak(let registerSuperNak): case .registerSuperNak(let registerSuperNak):
self.handleRegisterSuperNak(nakPacket: registerSuperNak) await self.handleRegisterSuperNak(nakPacket: registerSuperNak)
case .peerInfo(let peerInfo): case .peerInfo(let peerInfo):
SDLLogger.log("[SDLContext] peer message: \(peerInfo)") SDLLogger.log("[SDLContext] peer message: \(peerInfo)")
let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo) let packets = await self.puncherActor.makeRegisterPackets(peerInfo: peerInfo)
for packet in packets { for packet in packets {
await self.udpHoleServiceProxy.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress) await self.udpHoleService.send(type: .register, data: packet.data, remoteAddress: packet.remoteAddress)
} }
case .event(let event): case .event(let event):
await self.handleEvent(event: event) await self.handleEvent(event: event)
@ -321,6 +405,7 @@ extension SDLContextActor {
guard let key = try? self.rsaCipher.decode(data: Data(registerSuperAck.key)) else { guard let key = try? self.rsaCipher.decode(data: Data(registerSuperAck.key)) else {
SDLLogger.log("[SDLContext] registerSuperAck invalid key") SDLLogger.log("[SDLContext] registerSuperAck invalid key")
let error = SDLError.invalidKey let error = SDLError.invalidKey
await self.readySignal.fail(error)
self.provider.cancelTunnelWithError(error) self.provider.cancelTunnelWithError(error)
return return
} }
@ -337,6 +422,7 @@ extension SDLContextActor {
default: default:
SDLLogger.log("[SDLContext] registerSuperAck invalid algorithm \(algorithm)") SDLLogger.log("[SDLContext] registerSuperAck invalid algorithm \(algorithm)")
let error = SDLError.unsupportedAlgorithm(algorithm: algorithm) let error = SDLError.unsupportedAlgorithm(algorithm: algorithm)
await self.readySignal.fail(error)
self.provider.cancelTunnelWithError(error) self.provider.cancelTunnelWithError(error)
return return
} }
@ -351,8 +437,10 @@ extension SDLContextActor {
await self.packetOutboundActor.startPacketReader() await self.packetOutboundActor.startPacketReader()
// //
await self.whenRegistedSuper() await self.whenRegistedSuper()
await self.readySignal.succeed(())
} catch let err { } catch let err {
SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)") SDLLogger.log("[SDLContext] setTunnelNetworkSettings get error: \(err)")
await self.readySignal.fail(err)
self.provider.cancelTunnelWithError(err) self.provider.cancelTunnelWithError(err)
} }
} }
@ -361,7 +449,7 @@ extension SDLContextActor {
private func whenRegistedSuper() async { private func whenRegistedSuper() async {
await self.updatePolicyWorker?.stop() await self.updatePolicyWorker?.stop()
let policyService = self.policyService let policyService = self.policyService
let superServiceProxy = self.superServiceProxy let superService = self.superService
let updatePolicyWorker = PeriodicWorker( let updatePolicyWorker = PeriodicWorker(
configuration: .init( configuration: .init(
@ -372,7 +460,7 @@ extension SDLContextActor {
), ),
operation: { operation: {
SDLLogger.log("[SDLContext] updatePolicyTask execute") SDLLogger.log("[SDLContext] updatePolicyTask execute")
await policyService.updatePolicy(superServiceProxy: superServiceProxy) await policyService.updatePolicy(superService: superService)
}, },
onError: { err in onError: { err in
SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)") SDLLogger.log("[SDLContext] updatePolicyTask stop with err: \(err)")
@ -385,7 +473,7 @@ extension SDLContextActor {
await self.startStunRequestTask() await self.startStunRequestTask()
} }
private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) { private func handleRegisterSuperNak(nakPacket: SDLRegisterSuperNak) async {
let errorMessage = nakPacket.errorMessage let errorMessage = nakPacket.errorMessage
guard let errorCode = SDLNAKErrorCode(rawValue: UInt8(nakPacket.errorCode)) else { guard let errorCode = SDLNAKErrorCode(rawValue: UInt8(nakPacket.errorCode)) else {
return return
@ -396,6 +484,7 @@ extension SDLContextActor {
self.publishTunnelEvent(code: Int(errorCode.rawValue), message: errorMessage) self.publishTunnelEvent(code: Int(errorCode.rawValue), message: errorMessage)
// 退 // 退
let error = NSError(domain: "com.jihe.punchnet.tun", code: -1) let error = NSError(domain: "com.jihe.punchnet.tun", code: -1)
await self.readySignal.fail(error)
self.provider.cancelTunnelWithError(error) self.provider.cancelTunnelWithError(error)
case .noIpAddress, .networkFault, .internalFault: case .noIpAddress, .networkFault, .internalFault:
@ -445,7 +534,7 @@ extension SDLContextActor {
if let registerSuperData = try? registerSuper.serializedData() { if let registerSuperData = try? registerSuper.serializedData() {
SDLLogger.log("[SDLContext] will send register super") SDLLogger.log("[SDLContext] will send register super")
await self.superServiceProxy.send(type: .registerSuper, data: registerSuperData) await self.superService.send(type: .registerSuper, data: registerSuperData)
} }
} }
@ -454,7 +543,7 @@ extension SDLContextActor {
return return
} }
await self.superServiceProxy.send(type: .exposedServiceRequest, data: requestData) await self.superService.send(type: .exposedServiceRequest, data: requestData)
} }
private func applyExposedServiceResponse(_ response: SDLExposedServiceResponse) async { private func applyExposedServiceResponse(_ response: SDLExposedServiceResponse) async {

View File

@ -42,7 +42,7 @@ actor PacketInboundActor {
private let policyService: PolicyService private let policyService: PolicyService
private let packetOutboundActor: PacketOutboundActor private let packetOutboundActor: PacketOutboundActor
private let arpResolver: ArpResolver private let arpResolver: ArpResolver
private let superServiceProxy: SDLSuperServiceProxy private let superService: SDLSuperService
private let flowTracer: SDLFlowTracer private let flowTracer: SDLFlowTracer
private var networkAddress: SDLConfiguration.NetworkAddress private var networkAddress: SDLConfiguration.NetworkAddress
@ -55,7 +55,7 @@ actor PacketInboundActor {
policyService: PolicyService, policyService: PolicyService,
packetOutboundActor: PacketOutboundActor, packetOutboundActor: PacketOutboundActor,
arpResolver: ArpResolver, arpResolver: ArpResolver,
superServiceProxy: SDLSuperServiceProxy, superService: SDLSuperService,
flowTracer: SDLFlowTracer) { flowTracer: SDLFlowTracer) {
self.provider = provider self.provider = provider
self.networkAddress = config.networkAddress self.networkAddress = config.networkAddress
@ -64,7 +64,7 @@ actor PacketInboundActor {
self.policyService = policyService self.policyService = policyService
self.packetOutboundActor = packetOutboundActor self.packetOutboundActor = packetOutboundActor
self.arpResolver = arpResolver self.arpResolver = arpResolver
self.superServiceProxy = superServiceProxy self.superService = superService
self.flowTracer = flowTracer self.flowTracer = flowTracer
} }
@ -96,7 +96,7 @@ actor PacketInboundActor {
case .requestPolicy(let context): case .requestPolicy(let context):
SDLLogger.log("[PacketInboundActor] policy miss, \(context.logDescription)", for: .debug) SDLLogger.log("[PacketInboundActor] policy miss, \(context.logDescription)", for: .debug)
if let queryData = await self.policyService.makePolicyRequest(srcIdentityID: context.srcIdentityID) { if let queryData = await self.policyService.makePolicyRequest(srcIdentityID: context.srcIdentityID) {
await self.superServiceProxy.send(type: .policyRequest, data: queryData) await self.superService.send(type: .policyRequest, data: queryData)
} }
case .dropByPolicy(let context): case .dropByPolicy(let context):
SDLLogger.log("[PacketInboundActor] policy denied, \(context.logDescription)", for: .trace) SDLLogger.log("[PacketInboundActor] policy denied, \(context.logDescription)", for: .trace)

View File

@ -23,8 +23,8 @@ actor PacketOutboundActor {
private let arpResolver: ArpResolver private let arpResolver: ArpResolver
private let puncherActor: SDLPuncherActor private let puncherActor: SDLPuncherActor
private let policyService: PolicyService private let policyService: PolicyService
private let superServiceProxy: SDLSuperServiceProxy private let superService: SDLSuperService
private let udpHoleServiceProxy: SDLUDPHoleServiceProxy private let udpHoleService: SDLUDPHoleService
private let flowTracer: SDLFlowTracer private let flowTracer: SDLFlowTracer
private var packetReaderTask: Task<Void, Never>? private var packetReaderTask: Task<Void, Never>?
private var packetReaderGeneration: UInt64 = 0 private var packetReaderGeneration: UInt64 = 0
@ -43,8 +43,8 @@ actor PacketOutboundActor {
arpResolver: ArpResolver, arpResolver: ArpResolver,
puncherActor: SDLPuncherActor, puncherActor: SDLPuncherActor,
policyService: PolicyService, policyService: PolicyService,
superServiceProxy: SDLSuperServiceProxy, superService: SDLSuperService,
udpHoleServiceProxy: SDLUDPHoleServiceProxy, udpHoleService: SDLUDPHoleService,
flowTracer: SDLFlowTracer) { flowTracer: SDLFlowTracer) {
self.provider = provider self.provider = provider
self.networkAddress = config.networkAddress self.networkAddress = config.networkAddress
@ -56,8 +56,8 @@ actor PacketOutboundActor {
self.arpResolver = arpResolver self.arpResolver = arpResolver
self.puncherActor = puncherActor self.puncherActor = puncherActor
self.policyService = policyService self.policyService = policyService
self.superServiceProxy = superServiceProxy self.superService = superService
self.udpHoleServiceProxy = udpHoleServiceProxy self.udpHoleService = udpHoleService
self.flowTracer = flowTracer self.flowTracer = flowTracer
} }
@ -159,7 +159,7 @@ actor PacketOutboundActor {
} else { } else {
SDLLogger.log("[PacketOutboundActor] dstIp: \(SDLUtil.int32ToIp(ip)) arp query not found, broadcast", for: .trace) SDLLogger.log("[PacketOutboundActor] dstIp: \(SDLUtil.int32ToIp(ip)) arp query not found, broadcast", for: .trace)
if let arpRequest = try? await self.arpResolver.makeArpRequest(targetIp: ip) { if let arpRequest = try? await self.arpResolver.makeArpRequest(targetIp: ip) {
await self.superServiceProxy.send(type: .arpRequest, data: arpRequest) await self.superService.send(type: .arpRequest, data: arpRequest)
} }
} }
} }
@ -183,7 +183,7 @@ actor PacketOutboundActor {
self.flowTracer.inc(num: payload.count, type: .forward) self.flowTracer.inc(num: payload.count, type: .forward)
if let queryData = await self.puncherActor.makeQueryInfoRequest(request: request) { if let queryData = await self.puncherActor.makeQueryInfoRequest(request: request) {
await self.superServiceProxy.send(type: .queryInfo, data: queryData) await self.superService.send(type: .queryInfo, data: queryData)
} }
} }
@ -265,6 +265,6 @@ actor PacketOutboundActor {
} }
private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { private func sendPacket(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
await self.udpHoleServiceProxy.send(type: type, data: data, remoteAddress: remoteAddress) await self.udpHoleService.send(type: type, data: data, remoteAddress: remoteAddress)
} }
} }

View File

@ -59,10 +59,10 @@ actor PolicyService {
return await self.policyRuleStore.makePolicyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.identityId) return await self.policyRuleStore.makePolicyRequest(srcIdentityId: srcIdentityID, dstIdentityId: self.identityId)
} }
func updatePolicy(superServiceProxy: SDLSuperServiceProxy) async { func updatePolicy(superService: SDLSuperService) async {
let requests = await self.policyRuleStore.makeBatchPolicyRequests(dstIdentityID: self.identityId) let requests = await self.policyRuleStore.makeBatchPolicyRequests(dstIdentityID: self.identityId)
for request in requests { for request in requests {
await superServiceProxy.send(type: .policyRequest, data: request) await superService.send(type: .policyRequest, data: request)
} }
} }

View File

@ -5,109 +5,160 @@ actor SDLSuperService {
private let serverEndpoint: SDLConfiguration.ResolvedServerEndpoint private let serverEndpoint: SDLConfiguration.ResolvedServerEndpoint
private let port: UInt16 private let port: UInt16
private let onMessage: MessageHandler
private var superClient: SDLSuperClient? private var onMessage: MessageHandler = { _ in }
private var monitorTask: Task<Void, Never>? private var currentSession: SDLSuperSession?
private var generation: UInt64 = 0
init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16 = 1443, onMessage: @escaping MessageHandler) { init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16 = 1443) {
self.serverEndpoint = serverEndpoint self.serverEndpoint = serverEndpoint
self.port = port self.port = port
}
func updateMessageHandler(_ onMessage: @escaping MessageHandler) {
self.onMessage = onMessage self.onMessage = onMessage
} }
func start() { func run() async throws {
guard self.monitorTask == nil else { let generation = self.nextGeneration()
return let session = SDLSuperSession(
} serverEndpoint: self.serverEndpoint,
port: self.port,
self.monitorTask = startMonitorTask(name: "superServiceMonitor") { [weak self] in onMessage: { [weak self] message in
guard let self else { await self?.handleMessage(message, generation: generation)
throw CancellationError()
} }
try await self.runOnce() )
}
}
func stop() async { self.currentSession = session
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(serverEndpoint: self.serverEndpoint, port: self.port)
self.superClient = superClient
await superClient.start()
do { do {
try await withTaskCancellationHandler { try await session.run()
try await self.run(superClient) self.clearCurrent(session, generation: generation)
} onCancel: { } catch is CancellationError {
Task { self.clearCurrent(session, generation: generation)
await superClient.stop() await session.stop()
} throw CancellationError()
}
await self.cleanup(superClient)
} catch { } catch {
await self.cleanup(superClient) self.clearCurrent(session, generation: generation)
await session.stop()
throw error throw error
} }
} }
private func run(_ superClient: SDLSuperClient) async throws { func stop() async {
self.generation &+= 1
let session = self.currentSession
self.currentSession = nil
await session?.stop()
}
func send(type: SDLPacketType, data: Data) async {
await self.currentSession?.send(type: type, data: data)
}
private func nextGeneration() -> UInt64 {
self.generation &+= 1
return self.generation
}
private func clearCurrent(_ session: SDLSuperSession, generation: UInt64) {
guard self.generation == generation else {
return
}
if self.currentSession === session {
self.currentSession = nil
}
}
private func handleMessage(_ message: SDLQUICInboundMessage, generation: UInt64) async {
guard self.generation == generation else {
return
}
await self.onMessage(message)
}
}
final class SDLSuperSession: @unchecked Sendable {
typealias MessageHandler = @Sendable (SDLQUICInboundMessage) async -> Void
private let serverEndpoint: SDLConfiguration.ResolvedServerEndpoint
private let port: UInt16
private let onMessage: MessageHandler
private let client: SDLSuperClient
init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16, onMessage: @escaping MessageHandler) {
self.serverEndpoint = serverEndpoint
self.port = port
self.onMessage = onMessage
self.client = SDLSuperClient(serverEndpoint: serverEndpoint, port: port)
}
func run() async throws {
await self.client.start()
do {
try await withTaskCancellationHandler {
try await self.runLoops()
} onCancel: {
Task {
await self.client.stop()
}
}
await self.stop()
} catch {
await self.stop()
throw error
}
}
func stop() async {
await self.client.stop()
}
func send(type: SDLPacketType, data: Data) async {
await self.client.send(type: type, data: data)
}
private func runLoops() async throws {
try await Task.sleep(for: .seconds(0.5)) try await Task.sleep(for: .seconds(0.5))
try Task.checkCancellation() try Task.checkCancellation()
SDLLogger.log("[SDLSuperService] start super client: \(self.serverEndpoint.ip)") SDLLogger.log("[SDLSuperSession] start super client: \(self.serverEndpoint.ip)")
try await withThrowingTaskGroup(of: Void.self) { group in try await withThrowingTaskGroup(of: Void.self) { group in
defer { defer {
group.cancelAll() group.cancelAll()
} }
let onMessage = self.onMessage
group.addTask { group.addTask {
for try await message in superClient.messageStream { try await self.readLoop()
try Task.checkCancellation()
await onMessage(message)
}
} }
group.addTask { group.addTask {
while true { try await self.pingLoop()
try await Task.sleep(for: .seconds(5))
try Task.checkCancellation()
await superClient.send(type: .ping, data: Data())
}
} }
_ = try await group.next() _ = try await group.next()
} }
} }
private func cleanup(_ superClient: SDLSuperClient) async { private func readLoop() async throws {
await superClient.stop() for try await message in self.client.messageStream {
try Task.checkCancellation()
if self.superClient === superClient { await self.onMessage(message)
self.superClient = nil
} }
}
SDLLogger.log("[SDLSuperService] cleanup") private func pingLoop() async throws {
while true {
try await Task.sleep(for: .seconds(5))
try Task.checkCancellation()
await self.client.send(type: .ping, data: Data())
}
} }
} }

View File

@ -27,27 +27,131 @@ actor SDLUDPHoleService {
typealias DataHandler = @Sendable (SDLData) async -> Void typealias DataHandler = @Sendable (SDLData) async -> Void
private let proberActor: SDLNATProberActor private let proberActor: SDLNATProberActor
private let onEvent: EventHandler
private let onData: DataHandler
private var udpHole: SDLUDPHole? private var onEvent: EventHandler = { _ in }
private var udpHoleMonitorTask: Task<Void, Never>? private var onData: DataHandler = { _ in }
private var natProbeTask: Task<Void, Never>? private var currentSession: SDLUDPHoleSession?
private var localAddress: SocketAddress? private var generation: UInt64 = 0
private var udpHoleV6: SDLUDPHoleV6? init(proberActor: SDLNATProberActor) {
private var udpHoleV6MonitorTask: Task<Void, Never>?
init(proberActor: SDLNATProberActor, onEvent: @escaping EventHandler, onData: @escaping DataHandler) {
self.proberActor = proberActor self.proberActor = proberActor
}
func updateHandlers(onEvent: @escaping EventHandler, onData: @escaping DataHandler) {
self.onEvent = onEvent self.onEvent = onEvent
self.onData = onData self.onData = onData
} }
func start(includeV6: Bool = false) { func run(includeV6: Bool = false) async throws {
self.startV4() let generation = self.nextGeneration()
if includeV6 { let session = SDLUDPHoleSession(
self.startV6() proberActor: self.proberActor,
includeV6: includeV6,
onEvent: { [weak self] event in
await self?.handleEvent(event, generation: generation)
},
onData: self.onData
)
self.currentSession = session
do {
try await session.run()
self.clearCurrent(session, generation: generation)
} catch is CancellationError {
self.clearCurrent(session, generation: generation)
await session.stop()
throw CancellationError()
} catch {
self.clearCurrent(session, generation: generation)
await session.stop()
throw error
}
}
func stop() async {
self.generation &+= 1
let session = self.currentSession
self.currentSession = nil
await session?.stop()
await self.proberActor.cancelAll()
}
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
await self.currentSession?.send(type: type, data: data, remoteAddress: remoteAddress)
}
private func nextGeneration() -> UInt64 {
self.generation &+= 1
return self.generation
}
private func clearCurrent(_ session: SDLUDPHoleSession, generation: UInt64) {
guard self.generation == generation else {
return
}
if self.currentSession === session {
self.currentSession = nil
}
}
private func handleEvent(_ event: Event, generation: UInt64) async {
guard self.generation == generation else {
return
}
await self.onEvent(event)
}
}
actor SDLUDPHoleSession {
private let proberActor: SDLNATProberActor
private let includeV6: Bool
private let onEvent: SDLUDPHoleService.EventHandler
private let onData: SDLUDPHoleService.DataHandler
private var udpHole: SDLUDPHole?
private var udpHoleV6: SDLUDPHoleV6?
private var localAddress: SocketAddress?
init(
proberActor: SDLNATProberActor,
includeV6: Bool,
onEvent: @escaping SDLUDPHoleService.EventHandler,
onData: @escaping SDLUDPHoleService.DataHandler
) {
self.proberActor = proberActor
self.includeV6 = includeV6
self.onEvent = onEvent
self.onData = onData
}
func run() async throws {
do {
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
group.addTask {
try await self.runV4()
}
if self.includeV6 {
group.addTask {
try await self.runV6()
}
}
try await group.waitForAll()
}
await self.stop()
} catch {
await self.stop()
throw error
} }
} }
@ -56,64 +160,32 @@ actor SDLUDPHoleService {
self.udpHole = nil self.udpHole = nil
self.localAddress = 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 let udpHoleV6 = self.udpHoleV6
self.udpHoleV6 = nil self.udpHoleV6 = nil
let udpHoleV6MonitorTask = self.udpHoleV6MonitorTask
self.udpHoleV6MonitorTask = nil await self.proberActor.cancelAll()
udpHoleV6MonitorTask?.cancel() await udpHole?.stop()
udpHoleV6?.stop() udpHoleV6?.stop()
if let udpHoleV6MonitorTask {
await udpHoleV6MonitorTask.value
}
} }
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async { func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) async {
switch remoteAddress { switch remoteAddress {
case .v4: case .v4:
guard let udpHole else { guard let udpHole else {
SDLLogger.log("[SDLUDPHoleService] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug) SDLLogger.log("[SDLUDPHoleSession] udpHole is nil for remoteAddress: \(remoteAddress)", for: .debug)
return return
} }
await udpHole.send(type: type, data: data, remoteAddress: remoteAddress) await udpHole.send(type: type, data: data, remoteAddress: remoteAddress)
case .v6: case .v6:
guard let udpHoleV6 else { guard let udpHoleV6 else {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug) SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 is nil for remoteAddress: \(remoteAddress)", for: .debug)
return return
} }
udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress) udpHoleV6.send(type: type, data: data, remoteAddress: remoteAddress)
default: default:
SDLLogger.log("[SDLUDPHoleService] unsupported socket family: \(remoteAddress)", for: .debug) SDLLogger.log("[SDLUDPHoleSession] 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()
} }
} }
@ -122,23 +194,26 @@ actor SDLUDPHoleService {
let localAddress = try await udpHole.start() let localAddress = try await udpHole.start()
self.udpHole = udpHole self.udpHole = udpHole
self.localAddress = localAddress self.localAddress = localAddress
SDLLogger.log("[SDLUDPHoleService] udpHole started, on address: \(localAddress)")
SDLLogger.log("[SDLUDPHoleSession] udpHole started, on address: \(localAddress)")
await self.onEvent(.ready(localAddress)) await self.onEvent(.ready(localAddress))
self.startNatProbe(using: udpHole)
defer {
if self.udpHole === udpHole {
self.udpHole = nil
self.localAddress = nil
}
}
do { do {
try await withTaskCancellationHandler { try await withTaskCancellationHandler {
for try await datagram in await udpHole.messageStream() { try await withThrowingTaskGroup(of: Void.self) { group in
try Task.checkCancellation() defer {
try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message) group.cancelAll()
}
group.addTask {
try await self.readV4Loop(udpHole: udpHole)
}
group.addTask {
await self.probeNatType(udpHole: udpHole)
}
try await group.waitForAll()
} }
} onCancel: { } onCancel: {
Task { Task {
@ -147,26 +222,34 @@ actor SDLUDPHoleService {
} }
} catch { } catch {
await udpHole.stop() await udpHole.stop()
if self.udpHole === udpHole {
self.udpHole = nil
self.localAddress = nil
}
throw error throw error
} }
} }
private func startNatProbe(using udpHole: SDLUDPHole) { private func readV4Loop(udpHole: SDLUDPHole) async throws {
self.natProbeTask?.cancel() for try await datagram in await udpHole.messageStream() {
let proberActor = self.proberActor try Task.checkCancellation()
let onEvent = self.onEvent try await self.handleV4Message(remoteAddress: datagram.remoteAddress, message: datagram.message)
self.natProbeTask = Task {
if Task.isCancelled {
return
}
let natType = await proberActor.probeNatType(using: udpHole)
if Task.isCancelled {
return
}
await onEvent(.natType(natType))
} }
} }
private func probeNatType(udpHole: SDLUDPHole) async {
if Task.isCancelled {
return
}
let natType = await self.proberActor.probeNatType(using: udpHole)
if Task.isCancelled {
return
}
await self.onEvent(.natType(natType))
}
private func handleV4Message(remoteAddress: SocketAddress, message: SDLHoleMessage) async throws { private func handleV4Message(remoteAddress: SocketAddress, message: SDLHoleMessage) async throws {
switch message { switch message {
case .control(let control): case .control(let control):
@ -181,69 +264,57 @@ actor SDLUDPHoleService {
} }
} }
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 { private func runV6() async throws {
let udpHoleV6 = try SDLUDPHoleV6() let udpHoleV6 = try SDLUDPHoleV6()
let localAddress = try udpHoleV6.start() let localAddress = try udpHoleV6.start()
self.udpHoleV6 = udpHoleV6 self.udpHoleV6 = udpHoleV6
if let localAddress { if let localAddress {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, on address: \(localAddress)") SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, on address: \(localAddress)")
} else { } else {
SDLLogger.log("[SDLUDPHoleService] udpHoleV6 started, no local address") SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 started, no local address")
} }
defer { do {
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
let onEvent = self.onEvent
let onData = self.onData
group.addTask {
for await (remoteAddress, message) in udpHoleV6.messageStream {
try Task.checkCancellation()
switch message {
case .control(let control):
await onEvent(.packet(remoteAddress, control, source: .v6))
case .data(let data):
await onData(data)
}
}
}
group.addTask {
for await event in udpHoleV6.eventStream {
try Task.checkCancellation()
switch event {
case .ready:
SDLLogger.log("[SDLUDPHoleSession] udpHoleV6 ready")
case .closed, .errorCaught:
throw SDLContextError.udpHoleClosed
}
}
}
_ = try await group.next()
}
} catch {
udpHoleV6.stop()
if self.udpHoleV6 === udpHoleV6 { if self.udpHoleV6 === udpHoleV6 {
udpHoleV6.stop()
self.udpHoleV6 = nil self.udpHoleV6 = nil
} }
} throw error
try await withThrowingTaskGroup(of: Void.self) { group in
defer {
group.cancelAll()
}
let onEvent = self.onEvent
let onData = self.onData
group.addTask {
for await (remoteAddress, message) in udpHoleV6.messageStream {
try Task.checkCancellation()
switch message {
case .control(let control):
await onEvent(.packet(remoteAddress, control, source: .v6))
case .data(let data):
await onData(data)
}
}
}
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()
} }
} }
} }