punchnet-macos/Tun/PacketTunnelProvider.swift

276 lines
8.2 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.

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