276 lines
8.2 KiB
Swift
276 lines
8.2 KiB
Swift
//
|
||
// PacketTunnelProvider.swift
|
||
// punchnet
|
||
//
|
||
// Created by 安礼成 on 2025/8/3.
|
||
//
|
||
|
||
import Foundation
|
||
import NetworkExtension
|
||
|
||
enum TunnelError: Error {
|
||
case invalidConfiguration
|
||
case invalidContext
|
||
}
|
||
|
||
class PacketTunnelProvider: NEPacketTunnelProvider {
|
||
private var runtimeEnv: SDLRuntimeEnvironment?
|
||
|
||
override func startTunnel(options: [String: NSObject]?, completionHandler: @escaping (Error?) -> Void) {
|
||
// 如果当前在运行状态,不允许重复请求
|
||
guard self.runtimeEnv == nil else {
|
||
completionHandler(TunnelError.invalidContext)
|
||
return
|
||
}
|
||
|
||
guard let options, let config = SDLConfiguration.parse(options: options) else {
|
||
completionHandler(TunnelError.invalidConfiguration)
|
||
return
|
||
}
|
||
|
||
let rsaCipher = try! CCRSACipher(keySize: 1024)
|
||
self.runtimeEnv = SDLRuntimeEnvironment(config: config, rsaCipher: rsaCipher)
|
||
Task {
|
||
do {
|
||
try await self.runtimeEnv?.start(provider: self)
|
||
completionHandler(nil)
|
||
} catch let err {
|
||
completionHandler(err)
|
||
}
|
||
}
|
||
|
||
}
|
||
|
||
override func stopTunnel(with reason: NEProviderStopReason, completionHandler: @escaping () -> Void) {
|
||
// Add code here to start the process of stopping the tunnel.
|
||
Task {
|
||
await self.runtimeEnv?.stop()
|
||
completionHandler()
|
||
}
|
||
}
|
||
|
||
override func handleAppMessage(_ messageData: Data, completionHandler: ((Data?) -> Void)?) {
|
||
// Add code here to handle the message.
|
||
Task {
|
||
do {
|
||
let message = try AppRequest(serializedBytes: messageData)
|
||
let replyData = try await self.handleAppRequest(message: message)
|
||
completionHandler?(replyData)
|
||
} catch let err {
|
||
var reply = TunnelResponse()
|
||
reply.code = 1
|
||
reply.message = err.localizedDescription
|
||
|
||
let errorReplyData = try? reply.serializedData()
|
||
completionHandler?(errorReplyData)
|
||
}
|
||
}
|
||
}
|
||
|
||
override func sleep(completionHandler: @escaping () -> Void) {
|
||
// Add code here to get ready to sleep.
|
||
Task {
|
||
await self.runtimeEnv?.stop()
|
||
completionHandler()
|
||
}
|
||
}
|
||
|
||
override func wake() {
|
||
SDLLogger.log("[PacketTunnelProvider] wake up!!!!!!!")
|
||
// 启动monitor
|
||
let monitor = SDLPathMonitor()
|
||
// 启动监视器,允许重入
|
||
monitor.start()
|
||
SDLLogger.log("[PacketTunnelProvider] monitor started")
|
||
|
||
// Add code here to wake up.
|
||
Task {
|
||
|
||
defer {
|
||
monitor.stop()
|
||
}
|
||
|
||
// 等待网络可达
|
||
_ = await monitor.statusStream().first {$0 == .satisfied}
|
||
SDLLogger.log("[PacketTunnelProvider] network is satisfied")
|
||
// 重新启动
|
||
try await self.runtimeEnv?.start(provider: self)
|
||
}
|
||
|
||
}
|
||
|
||
private func handleAppRequest(message: AppRequest) async throws -> Data? {
|
||
guard let contextActor = self.runtimeEnv?.getContextActor() else {
|
||
throw TunnelError.invalidContext
|
||
}
|
||
|
||
switch message.command {
|
||
case .changeExitNode(let changeExitNode):
|
||
let exitNodeIp = changeExitNode.ip
|
||
do {
|
||
try await contextActor.updateExitNode(exitNodeIp: exitNodeIp)
|
||
var reply = TunnelResponse()
|
||
reply.code = 0
|
||
reply.message = "操作成功"
|
||
return try reply.serializedData()
|
||
|
||
} catch let err {
|
||
var reply = TunnelResponse()
|
||
reply.code = 1
|
||
reply.message = err.localizedDescription
|
||
|
||
return try reply.serializedData()
|
||
}
|
||
case .none:
|
||
var reply = TunnelResponse()
|
||
reply.code = 1
|
||
reply.message = "无效请求"
|
||
return try reply.serializedData()
|
||
}
|
||
}
|
||
|
||
}
|
||
|
||
private class SDLRuntimeEnvironment {
|
||
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()
|
||
|
||
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.withStateLock {
|
||
self.contextActor
|
||
}
|
||
}
|
||
|
||
func stop() async {
|
||
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<T>(_ body: () -> T) -> T {
|
||
self.stateLock.lock()
|
||
defer {
|
||
self.stateLock.unlock()
|
||
}
|
||
return body()
|
||
}
|
||
|
||
}
|