punchnet-macos/Tun/Punchnet/Actors/SDLQuicClient.swift
2026-04-28 16:13:31 +08:00

479 lines
15 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.

//
// SDLQuicClient.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 timeout
case decodeError(String)
case packetTooLarge
}
enum SDLQUICEvent: Error {
case failed(Error)
case cancelled
case writeFailed(Error)
}
enum SDLQUICClientExit: Error, Sendable, CustomStringConvertible {
case normal
case cancelled
case transportClosed(String)
case readFailed(String)
case writeFailed(String)
var description: String {
switch self {
case .normal:
return "normal"
case .cancelled:
return "cancelled"
case .transportClosed(let reason):
return "transportClosed(\(reason))"
case .readFailed(let reason):
return "readFailed(\(reason))"
case .writeFailed(let reason):
return "writeFailed(\(reason))"
}
}
}
actor SDLQUICClient {
private let allocator = ByteBufferAllocator()
// 64K
private let maxPacketSize: Int
// 2M
private let maxBufferSize: Int
private let readyState = SDLQUICReadyState()
//
public var messageStream: AsyncStream<SDLQUICInboundMessage>
private let messageCont: AsyncStream<SDLQUICInboundMessage>.Continuation
//
public var eventStream: AsyncStream<SDLQUICEvent>
private let eventCont: AsyncStream<SDLQUICEvent>.Continuation
private let connection: NWConnection
private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 线
init(host: String, port: UInt16, maxPacketSize: Int = 64 * 1024, maxBufferSize: Int = 2 * 1024 * 1024) {
let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"])
self.maxBufferSize = maxBufferSize
self.maxPacketSize = maxPacketSize
(self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self)
(self.eventStream, self.eventCont) = AsyncStream.makeStream(of: SDLQUICEvent.self)
// TODO
sec_protocol_options_set_verify_block(
options.securityProtocolOptions,
{ metadata, trust, complete in
//
complete(QUICVerifier.verify(trust: trust, host: host))
},
self.queue
)
let params = NWParameters(quic: options)
self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
}
func start() {
connection.stateUpdateHandler = { state in
SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug)
switch state {
case .ready:
Task {
await self.readyState.markReady()
}
case .failed(let error):
Task {
await self.readyState.markFailed(error)
}
self.eventCont.yield(.failed(error))
case .cancelled:
Task {
await self.readyState.markCancelled()
}
self.eventCont.yield(.cancelled)
case .setup, .preparing:
Task {
await self.readyState.markConnecting()
}
default:
()
}
}
connection.start(queue: self.queue)
}
func waitReady(timeout: Duration = .seconds(5)) async throws {
try await withThrowingTaskGroup(of: Void.self) { group in
group.addTask {
try await self.readyState.waitReady()
}
group.addTask {
try await Task.sleep(for: timeout)
throw SDLQUICError.timeout
}
try await group.next()
group.cancelAll()
}
}
func run() async -> SDLQUICClientExit {
await withTaskGroup(of: SDLQUICClientExit.self) { group in
group.addTask {
await self.readLoop()
}
group.addTask {
await self.heartbeatLoop()
}
let exit = await group.next() ?? .normal
group.cancelAll()
return exit
}
}
func send(type: SDLPacketType, data: Data) {
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 {
SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug)
self?.eventCont.yield(.writeFailed(error))
}
})
}
private func heartbeatLoop() async -> SDLQUICClientExit {
let timerStream = SDLAsyncTimerStream()
timerStream.start(interval: .seconds(5))
for await _ in timerStream.stream {
if Task.isCancelled {
break
}
self.send(type: .ping, data: Data())
}
SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug)
return .cancelled
}
func stop() {
self.connection.cancel()
}
func close(_ exit: SDLQUICClientExit = .normal) async {
}
deinit {
self.messageCont.finish()
}
}
// --MARK: Ready
extension SDLQUICClient {
actor SDLQUICReadyState {
enum State {
case idle
case connecting
case ready
case failed(Error)
case cancelled
}
private var state: State = .idle
private var continuations: [CheckedContinuation<Void, Error>] = []
func waitReady() async throws {
switch state {
case .ready:
return
case .failed(let error):
throw error
case .cancelled:
throw CancellationError()
case .idle, .connecting:
try await withCheckedThrowingContinuation { continuation in
continuations.append(continuation)
}
}
}
func markConnecting() {
switch state {
case .idle:
state = .connecting
default:
break
}
}
func markReady() {
state = .ready
let list = continuations
continuations.removeAll()
for continuation in list {
continuation.resume()
}
}
func markFailed(_ error: Error) {
state = .failed(error)
let list = continuations
continuations.removeAll()
for continuation in list {
continuation.resume(throwing: error)
}
}
func markCancelled() {
state = .cancelled
let list = continuations
continuations.removeAll()
for continuation in list {
continuation.resume(throwing: CancellationError())
}
}
}
}
// --MARK: Reader
extension SDLQUICClient {
private func readLoop() async -> SDLQUICClientExit {
var buffer = allocator.buffer(capacity: self.maxBufferSize)
let threshold = self.maxBufferSize / 10 * 6
defer {
self.messageCont.finish()
}
do {
while !Task.isCancelled {
let (isComplete, data) = try await self.readOnce()
if let data, !data.isEmpty {
buffer.writeBytes(data)
let frames = try parseFrames(buffer: &buffer)
if buffer.readerIndex > threshold {
buffer.discardReadBytes()
}
for frame in frames {
if let message = decode(frame: frame) {
self.messageCont.yield(message)
}
}
}
if isComplete {
return .transportClosed("receive complete")
}
}
return .cancelled
} catch is CancellationError {
return .cancelled
} catch {
return .readFailed("\(error)")
}
}
//
private func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] {
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)
}
}
return frames
}
//
private func readOnce() async throws -> (Bool, Data?) {
return try await withCheckedThrowingContinuation { cont in
self.connection.receive(minimumIncompleteLength: 1, maximumLength: maxPacketSize) { data, _, isComplete, error in
if let error {
cont.resume(throwing: error)
return
}
cont.resume(returning: (isComplete, data))
}
}
}
}
// --MARK:
extension SDLQUICClient {
private 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("SDLQUICClient decode Event Error", for: .debug)
return nil
}
return .event(event)
case .pong:
return .pong
default:
SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug)
return nil
}
}
}
// --MARK: quic
extension SDLQUICClient {
enum QUICVerifier {
// 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
}
}
}
}