diff --git a/Tun/PacketTunnelProvider.swift b/Tun/PacketTunnelProvider.swift index 7db99d7..0f81428 100644 --- a/Tun/PacketTunnelProvider.swift +++ b/Tun/PacketTunnelProvider.swift @@ -5,6 +5,7 @@ // Created by 安礼成 on 2025/8/3. // +import Foundation import NetworkExtension enum TunnelError: Error { @@ -22,12 +23,13 @@ class PacketTunnelProvider: NEPacketTunnelProvider { return } - guard let options else { + guard let options, let config = SDLConfiguration.parse(options: options) else { completionHandler(TunnelError.invalidConfiguration) return } - self.runtimeEnv = SDLRuntimeEnvironment(options: options) + let rsaCipher = try! CCRSACipher(keySize: 1024) + self.runtimeEnv = SDLRuntimeEnvironment(config: config, rsaCipher: rsaCipher) Task { do { try await self.runtimeEnv?.start(provider: self) @@ -130,35 +132,144 @@ class PacketTunnelProvider: NEPacketTunnelProvider { } private class SDLRuntimeEnvironment { - var contextActor: SDLContextActor? - private var options: [String: NSObject] - - init(options: [String: NSObject]) { - self.options = options + private enum State { + case idle + case starting + case running + case stopping + } + + private let stateLock = NSLock() + private var state: State = .idle + private var pendingStop = false + private weak var pendingStartProvider: PacketTunnelProvider? + private var contextActor: SDLContextActor? + private var config: SDLConfiguration + private let rsaCipher: CCRSACipher + + init(config: SDLConfiguration, rsaCipher: CCRSACipher) { + self.config = config + self.rsaCipher = rsaCipher } func start(provider: PacketTunnelProvider) async throws { + guard self.markStarting(provider: provider) else { + return + } + // 重置通知中心 SDLTunnelAppNotifier.shared.clear() - guard let config = await SDLConfiguration.parse(options: options) else { - throw TunnelError.invalidConfiguration - } - - // 加密算法 - let rsaCipher = try! CCRSACipher(keySize: 1024) - let contextActor = SDLContextActor(provider: provider, config: config, rsaCipher: rsaCipher) - self.contextActor = contextActor + let contextActor = SDLContextActor(provider: provider, config: config, rsaCipher: self.rsaCipher) await contextActor.start() + + let shouldStop = self.markStarted(contextActor: contextActor) + if shouldStop { + await self.stop() + } } func getContextActor() -> SDLContextActor? { - return self.contextActor + return self.withStateLock { + self.contextActor + } } func stop() async { - await self.contextActor?.stop() - self.contextActor = nil + guard let contextActor = self.markStopping() else { + return + } + + await contextActor.stop() + + if let provider = self.markStopped() { + try? await self.start(provider: provider) + } + } + + private func markStarting(provider: PacketTunnelProvider) -> Bool { + return self.withStateLock { + switch self.state { + case .idle: + self.pendingStop = false + self.pendingStartProvider = nil + self.state = .starting + return true + + case .starting, .running: + SDLLogger.log("[SDLRuntimeEnvironment] skip duplicated start: \(self.state)", for: .debug) + return false + + case .stopping: + self.pendingStartProvider = provider + SDLLogger.log("[SDLRuntimeEnvironment] delay start until stop finishes", for: .debug) + return false + } + } + } + + private func markStarted(contextActor: SDLContextActor) -> Bool { + return self.withStateLock { + self.contextActor = contextActor + self.state = .running + + let shouldStop = self.pendingStop + self.pendingStop = false + return shouldStop + } + } + + private func markStartFailed() { + self.withStateLock { + self.contextActor = nil + self.pendingStop = false + self.pendingStartProvider = nil + self.state = .idle + } + } + + private func markStopping() -> SDLContextActor? { + return self.withStateLock { + switch self.state { + case .idle: + self.contextActor = nil + self.pendingStop = false + return nil + + case .starting: + self.pendingStop = true + SDLLogger.log("[SDLRuntimeEnvironment] delay stop until start finishes", for: .debug) + return nil + + case .stopping: + SDLLogger.log("[SDLRuntimeEnvironment] skip duplicated stop", for: .debug) + return nil + + case .running: + self.state = .stopping + let contextActor = self.contextActor + self.contextActor = nil + self.pendingStop = false + return contextActor + } + } + } + + private func markStopped() -> PacketTunnelProvider? { + return self.withStateLock { + self.state = .idle + let provider = self.pendingStartProvider + self.pendingStartProvider = nil + return provider + } + } + + private func withStateLock(_ body: () -> T) -> T { + self.stateLock.lock() + defer { + self.stateLock.unlock() + } + return body() } } diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index 7692d88..3dc65ad 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -323,7 +323,7 @@ actor SDLContextActor { self.dnsWorker = nil // 启动dns服务 - let dnsClient = DNSCloudClient(host: self.config.serverIp, port: 15353) + let dnsClient = DNSCloudClient(host: self.config.serverHost, port: 15353) await dnsClient.start() SDLLogger.log("[SDLContext] dnsClient started") self.dnsClient = dnsClient diff --git a/Tun/Punchnet/SDLConfiguration.swift b/Tun/Punchnet/SDLConfiguration.swift index 17c3d1f..6f7be84 100644 --- a/Tun/Punchnet/SDLConfiguration.swift +++ b/Tun/Punchnet/SDLConfiguration.swift @@ -49,7 +49,6 @@ public class SDLConfiguration { let version: Int let serverHost: String - let serverIp: String let stunServers: [String] lazy var stunSocketAddress: SocketAddress = { @@ -77,7 +76,6 @@ public class SDLConfiguration { public init(version: Int, serverHost: String, - serverIp: String, stunServers: [String], clientId: String, networkAddress: NetworkAddress, @@ -87,7 +85,6 @@ public class SDLConfiguration { exitNode: ExitNode?) { self.version = version self.serverHost = serverHost - self.serverIp = serverIp self.stunServers = stunServers self.clientId = clientId self.networkAddress = networkAddress @@ -102,7 +99,7 @@ public class SDLConfiguration { // 解析配置文件 extension SDLConfiguration { - static func parse(options: [String: NSObject]) async -> SDLConfiguration? { + static func parse(options: [String: NSObject]) -> SDLConfiguration? { guard let version = options["version"] as? Int, let serverHost = options["server_host"] as? String, let stunAssistHost = options["stun_assist_host"] as? String, @@ -118,11 +115,6 @@ extension SDLConfiguration { return nil } - // 解析dns域名所在的服务器地址 - guard let serverIp = await SDLUtil.resolveHostname(host: serverHost) else { - return nil - } - // 网络出口配置是可选的 var exitNode: ExitNode? = nil if let exitNodeIpStr = options["exit_node_ip"] as? String, let exitNodeIp = SDLUtil.ipv4StrToInt32(exitNodeIpStr) { @@ -131,7 +123,6 @@ extension SDLConfiguration { return SDLConfiguration(version: version, serverHost: serverHost, - serverIp: serverIp, stunServers: [serverHost, stunAssistHost], clientId: clientId, networkAddress: networkAddress, diff --git a/punchnet/Features/Network/ViewModels/NetworkModel.swift b/punchnet/Features/Network/ViewModels/NetworkModel.swift index 5dddad7..5ef029e 100644 --- a/punchnet/Features/Network/ViewModels/NetworkModel.swift +++ b/punchnet/Features/Network/ViewModels/NetworkModel.swift @@ -7,6 +7,7 @@ import Foundation import Observation +@MainActor @Observable final class NetworkModel { @ObservationIgnored