407 lines
14 KiB
Swift
407 lines
14 KiB
Swift
//
|
||
// 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 serverEndpoint: SDLConfiguration.ResolvedServerEndpoint
|
||
private let port: UInt16
|
||
|
||
init(serverEndpoint: SDLConfiguration.ResolvedServerEndpoint, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) {
|
||
self.serverEndpoint = serverEndpoint
|
||
self.port = port
|
||
|
||
self.frameParser = SDLQUICFrameParser(maxBufferSize: maxBufferSize)
|
||
(self.messageStream, self.messageCont) = AsyncThrowingStream.makeStream(of: SDLQUICInboundMessage.self)
|
||
}
|
||
|
||
func start() {
|
||
let serverEndpoint = self.serverEndpoint
|
||
let queue = DispatchQueue(label: "com.sdl.SuperClient.queue") // 专用队列保证线程安全
|
||
|
||
let options = NWProtocolTLS.Options()
|
||
serverEndpoint.host.withCString {
|
||
sec_protocol_options_set_tls_server_name(options.securityProtocolOptions, $0)
|
||
}
|
||
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: serverEndpoint.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: Self.makeEndpointHost(address: serverEndpoint.ip), 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 static func makeEndpointHost(address ip: String) -> NWEndpoint.Host {
|
||
if let ipv4Address = IPv4Address(ip) {
|
||
return .ipv4(ipv4Address)
|
||
}
|
||
|
||
if let ipv6Address = IPv6Address(ip) {
|
||
return .ipv6(ipv6Address)
|
||
}
|
||
|
||
preconditionFailure("invalid super server IP: \(ip)")
|
||
}
|
||
|
||
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
|
||
}
|
||
}
|
||
|
||
}
|
||
}
|