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

522 lines
16 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 var didFinishStreams = false
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)
//
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 = { [weak self] state in
Task {
await self?.handleConnectionStateUpdate(state)
}
}
connection.start(queue: self.queue)
}
private func handleConnectionStateUpdate(_ state: NWConnection.State) async {
SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug)
switch state {
case .ready:
await self.readyState.markReady()
case .failed(let error):
await self.readyState.markFailed(error)
self.emitEvent(.failed(error))
case .cancelled:
await self.readyState.markCancelled()
self.emitEvent(.cancelled)
case .setup, .preparing:
await self.readyState.markConnecting()
default:
()
}
}
func waitReady(timeout: Duration = .seconds(5)) async throws {
try await self.readyState.waitReady(timeout: timeout)
}
func run() async -> SDLQUICClientExit {
await withTaskCancellationHandler {
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()
await self.stop()
self.finishStreams()
return exit
}
} onCancel: {
Task {
await self.stop()
}
}
}
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)
Task {
await self?.emitEvent(.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() async {
self.connection.cancel()
await self.readyState.markCancelled()
self.finishStreams()
}
private func emitEvent(_ event: SDLQUICEvent) {
guard !self.didFinishStreams else {
return
}
self.eventCont.yield(event)
}
private func finishStreams() {
guard !self.didFinishStreams else {
return
}
self.didFinishStreams = true
self.messageCont.finish()
self.eventCont.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: [UUID: CheckedContinuation<Void, Error>] = [:]
func waitReady(timeout: Duration) async throws {
let id = UUID()
let timeoutTask = Task {
try? await Task.sleep(for: timeout)
if Task.isCancelled {
return
}
self.cancelWaiter(id: id, throwing: SDLQUICError.timeout)
}
defer {
timeoutTask.cancel()
}
try await withTaskCancellationHandler {
try await withCheckedThrowingContinuation { continuation in
self.addWaiter(id: id, continuation: continuation)
}
} onCancel: {
timeoutTask.cancel()
Task {
await self.cancelWaiter(id: id, throwing: CancellationError())
}
}
}
private func addWaiter(id: UUID, continuation: CheckedContinuation<Void, Error>) {
switch state {
case .ready:
continuation.resume()
case .failed(let error):
continuation.resume(throwing: error)
case .cancelled:
continuation.resume(throwing: CancellationError())
case .idle, .connecting:
continuations[id] = continuation
}
}
private func cancelWaiter(id: UUID, throwing error: Error) {
guard let continuation = continuations.removeValue(forKey: id) else {
return
}
continuation.resume(throwing: error)
}
func markConnecting() {
switch state {
case .idle:
state = .connecting
default:
break
}
}
func markReady() {
state = .ready
let list = continuations
continuations.removeAll()
for continuation in list.values {
continuation.resume()
}
}
func markFailed(_ error: Error) {
state = .failed(error)
let list = continuations
continuations.removeAll()
for continuation in list.values {
continuation.resume(throwing: error)
}
}
func markCancelled() {
state = .cancelled
let list = continuations
continuations.removeAll()
for continuation in list.values {
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
}
}
}
}