punchnet-macos/Tun/Punchnet/Super/SDLSuperClient.swift
2026-05-07 10:34:43 +08:00

392 lines
13 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.

//
// SDLSuperClient.swift
// Tun
//
// Created by on 2026/2/13.
//
import Foundation
import NIOCore
import Network
import CryptoKit
import Security
// 便
enum SDLQUICError: Error {
case connectionFailed(Error)
case connectionCancelled
case writeFailed(Error)
case internalError(Error)
case timeout
case decodeError(String)
case packetTooLarge
case dataStreamClosed
}
actor SDLSuperClient {
enum State {
case idle
case running
case stopped
}
private var state: State = .idle
private let frameParser: SDLQUICFrameParser
//
public var messageStream: AsyncThrowingStream<SDLQUICInboundMessage, Error>
private let messageCont: AsyncThrowingStream<SDLQUICInboundMessage, Error>.Continuation
private var isMessageContinuationFinished: Bool = false
private var readTask: Task<Void, Never>?
private var connection: NWConnection?
private let host: String
private let port: UInt16
init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) {
self.host = host
self.port = port
self.frameParser = SDLQUICFrameParser(maxBufferSize: maxBufferSize)
(self.messageStream, self.messageCont) = AsyncThrowingStream.makeStream(of: SDLQUICInboundMessage.self)
}
func start() {
let host = self.host
let queue = DispatchQueue(label: "com.sdl.SuperClient.queue") // 线
let options = NWProtocolTLS.Options()
sec_protocol_options_add_tls_application_protocol(
options.securityProtocolOptions,
"punchnet/1.0"
)
//
sec_protocol_options_set_verify_block(
options.securityProtocolOptions,
{ _, trust, complete in
//
complete(TLSVerifier.verify(trust: trust, host: host))
},
queue
)
SDLLogger.log("[SDLSuperClient] start with tls protocol", for: .debug)
let params = NWParameters(tls: options)
// Network.framework
params.preferNoProxies = true
let connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
connection.stateUpdateHandler = { [weak self] state in
SDLLogger.log("[SDLSuperClient] new state: \(state)", for: .debug)
Task {
await self?.handleConnectionState(state: state)
}
}
connection.start(queue: queue)
self.connection = connection
}
private func handleConnectionState(state: NWConnection.State) {
switch state {
case .ready:
self.startReadTask()
self.state = .running
case .failed(let error):
self.finishMessageContinuationIfNeed(throwing: .connectionFailed(error))
case .cancelled:
self.finishMessageContinuationIfNeed(throwing: .connectionCancelled)
default:
()
}
}
private func finishMessageContinuationIfNeed(throwing error: SDLQUICError?) {
guard !self.isMessageContinuationFinished else {
return
}
self.isMessageContinuationFinished = true
if let error {
self.messageCont.finish(throwing: error)
} else {
self.messageCont.finish()
}
}
private func startReadTask() {
self.readTask?.cancel()
self.readTask = Task {
do {
while true {
try Task.checkCancellation()
let data = try await self.readOnce()
let frames = try self.frameParser.parseFrames(data: data)
for frame in frames {
if let message = SDLQUICCodec.decode(frame: frame) {
self.messageCont.yield(message)
} else {
self.finishMessageContinuationIfNeed(throwing: .decodeError("invalid message"))
}
}
}
} catch let err {
self.finishMessageContinuationIfNeed(throwing: .internalError(err))
}
}
}
func send(type: SDLPacketType, data: Data) {
guard case .running = state, let connection = self.connection, connection.state == .ready else {
return
}
var len = UInt16(data.count + 1).bigEndian
var packet = Data(Data(bytes: &len, count: 2))
packet.append(type.rawValue)
packet.append(data)
connection.send(content: packet, completion: .contentProcessed { [weak self] error in
if let error {
Task {
SDLLogger.log("[SDLSuperClient] send data get error: \(error)", for: .debug)
await self?.finishMessageContinuationIfNeed(throwing: .writeFailed(error))
}
}
})
}
private func readOnce() async throws -> Data {
guard let connection = self.connection else {
throw SDLQUICError.connectionCancelled
}
return try await withCheckedThrowingContinuation { cont in
connection.receive(minimumIncompleteLength: 1, maximumLength: 64 * 1024) { data, _, isComplete, error in
if let error {
cont.resume(throwing: error)
return
}
if isComplete {
cont.resume(throwing: SDLQUICError.dataStreamClosed)
} else {
cont.resume(returning: data ?? Data())
}
}
}
}
func stop() {
guard self.state != .stopped else {
return
}
self.state = .stopped
self.readTask?.cancel()
self.readTask = nil
let connection = self.connection
self.connection = nil
connection?.stateUpdateHandler = nil
connection?.cancel()
self.finishMessageContinuationIfNeed(throwing: nil)
SDLLogger.log("[SDLSuperClient] stopped")
}
deinit {
SDLLogger.log("[SDLSuperClient] deinit")
}
}
// --MARK:
extension SDLSuperClient {
class SDLQUICFrameParser {
private let allocator = ByteBufferAllocator()
// 2M
private let maxPacketSize: Int = 64 * 1024
private let maxBufferSize: Int
private var buffer: ByteBuffer
init(maxBufferSize: Int) {
self.buffer = allocator.buffer(capacity: maxBufferSize)
self.maxBufferSize = maxBufferSize
}
//
public func parseFrames(data: Data) throws -> [ByteBuffer] {
self.buffer.writeBytes(data)
guard buffer.readableBytes >= 2 else {
return []
}
var frames: [ByteBuffer] = []
while true {
guard let len = buffer.getInteger(at: buffer.readerIndex, endianness: .big, as: UInt16.self) else {
break
}
if len > self.maxPacketSize {
throw SDLQUICError.packetTooLarge
}
guard buffer.readableBytes >= len + 2 else {
break
}
buffer.moveReaderIndex(forwardBy: 2)
if let buf = buffer.readSlice(length: Int(len)) {
frames.append(buf)
}
}
let threshold = maxBufferSize / 10 * 6
if buffer.readerIndex > threshold {
buffer.discardReadBytes()
}
return frames
}
}
}
// --MARK:
extension SDLSuperClient {
enum SDLQUICCodec {
public static func decode(frame: ByteBuffer) -> SDLQUICInboundMessage? {
var buffer = frame
guard let type = buffer.readInteger(as: UInt8.self),
let packetType = SDLPacketType(rawValue: type) else {
return nil
}
switch packetType {
case .welcome:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let welcome = try? SDLWelcome(serializedBytes: bytes) else {
return nil
}
return .welcome(welcome)
case .registerSuperAck:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else {
return nil
}
return .registerSuperAck(registerSuperAck)
case .registerSuperNak:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let registerSuperNak = try? SDLRegisterSuperNak(serializedBytes: bytes) else {
return nil
}
return .registerSuperNak(registerSuperNak)
case .peerInfo:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let peerInfo = try? SDLPeerInfo(serializedBytes: bytes) else {
return nil
}
return .peerInfo(peerInfo)
case .policyResponse:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let policyResponse = try? SDLPolicyResponse(serializedBytes: bytes) else {
return nil
}
return .policyReponse(policyResponse)
case .arpResponse:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let arpResponse = try? SDLArpResponse(serializedBytes: bytes) else {
return nil
}
return .arpResponse(arpResponse)
case .event:
guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let event = try? SDLEvent(serializedBytes: bytes) else {
SDLLogger.log("SDLSuperClient decode Event Error", for: .debug)
return nil
}
return .event(event)
case .pong:
return .pong
default:
SDLLogger.log("SDLSuperClient decode miss type: \(type)", for: .debug)
return nil
}
}
}
}
// --MARK: tls
extension SDLSuperClient {
enum TLSVerifier {
// Base64
static let pinnedPublicKeyHashes = [
"Q41r6hbMWEVyxo6heNAH4Wx/TH5NNOWlNif9bewcJ3E="
]
static func verify(trust: sec_trust_t, host: String) -> Bool {
let secTrust = sec_trust_copy_ref(trust).takeRetainedValue()
// --- Step 1: ---
var error: CFError?
guard SecTrustEvaluateWithError(secTrust, &error) else {
SDLLogger.log("❌ 系统证书验证失败: \(error?.localizedDescription ?? "未知错误")", for: .debug)
return false
}
// --- Step 2: ---
let policy = SecPolicyCreateSSL(true, host as CFString)
SecTrustSetPolicies(secTrust, policy)
guard SecTrustEvaluateWithError(secTrust, &error) else {
SDLLogger.log("❌ 主机名校验失败: \(error?.localizedDescription ?? "未知错误")", for: .debug)
return false
}
// --- Step 3: ---
guard let chain = SecTrustCopyCertificateChain(secTrust) as? [SecCertificate],
let leafCertificate = chain.first else {
SDLLogger.log("❌ 无法获取证书链或叶子证书", for: .debug)
return false
}
// --- Step 4: ---
guard let publicKey = SecCertificateCopyKey(leafCertificate),
let publicKeyData = SecKeyCopyExternalRepresentation(publicKey, nil) as Data? else {
SDLLogger.log("❌ 无法提取公钥", for: .debug)
return false
}
// --- Step 5: SHA256 ---
let hash = SHA256.hash(data: publicKeyData)
let hashBase64 = Data(hash).base64EncodedString()
if pinnedPublicKeyHashes.contains(hashBase64) {
SDLLogger.log("✅ 公钥校验通过", for: .debug)
return true
} else {
SDLLogger.log("⚠️ 公钥不匹配! 收到: \(hashBase64)", for: .debug)
return false
}
}
}
}