fix quicClient

This commit is contained in:
anlicheng 2026-04-28 16:58:34 +08:00
parent b37bebd7b2
commit 7b549981af

View File

@ -18,6 +18,7 @@ enum SDLQUICError: Error {
case timeout case timeout
case decodeError(String) case decodeError(String)
case packetTooLarge case packetTooLarge
case waitReadyAlreadyInProgress
} }
enum SDLQUICEvent: Error { enum SDLQUICEvent: Error {
@ -32,7 +33,7 @@ enum SDLQUICClientExit: Error, Sendable, CustomStringConvertible {
case transportClosed(String) case transportClosed(String)
case readFailed(String) case readFailed(String)
case writeFailed(String) case writeFailed(String)
var description: String { var description: String {
switch self { switch self {
case .normal: case .normal:
@ -55,25 +56,35 @@ actor SDLQUICClient {
private let maxPacketSize: Int private let maxPacketSize: Int
// 2M // 2M
private let maxBufferSize: Int private let maxBufferSize: Int
private let readyState = SDLQUICReadyState() private enum ReadyStatus {
case idle
case connecting
case ready
case failed(Error)
case cancelled
}
private var readyStatus: ReadyStatus = .idle
private var readyContinuation: CheckedContinuation<Void, Error>?
private var readyTimeoutTask: Task<Void, Never>?
// //
public var messageStream: AsyncStream<SDLQUICInboundMessage> public var messageStream: AsyncStream<SDLQUICInboundMessage>
private let messageCont: AsyncStream<SDLQUICInboundMessage>.Continuation private let messageCont: AsyncStream<SDLQUICInboundMessage>.Continuation
// //
public var eventStream: AsyncStream<SDLQUICEvent> public var eventStream: AsyncStream<SDLQUICEvent>
private let eventCont: AsyncStream<SDLQUICEvent>.Continuation private let eventCont: AsyncStream<SDLQUICEvent>.Continuation
private var didFinishStreams = false private var didFinishStreams = false
private let connection: NWConnection private let connection: NWConnection
private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 线 private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 线
init(host: String, port: UInt16, maxPacketSize: Int = 64 * 1024, maxBufferSize: Int = 2 * 1024 * 1024) { init(host: String, port: UInt16, maxPacketSize: Int = 64 * 1024, maxBufferSize: Int = 2 * 1024 * 1024) {
let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"]) let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"])
self.maxBufferSize = maxBufferSize self.maxBufferSize = maxBufferSize
self.maxPacketSize = maxPacketSize self.maxPacketSize = maxPacketSize
(self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self) (self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self)
@ -88,7 +99,7 @@ actor SDLQUICClient {
}, },
self.queue self.queue
) )
let params = NWParameters(quic: options) let params = NWParameters(quic: options)
self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
} }
@ -101,27 +112,27 @@ actor SDLQUICClient {
} }
connection.start(queue: self.queue) connection.start(queue: self.queue)
} }
private func handleConnectionStateUpdate(_ state: NWConnection.State) async { private func handleConnectionStateUpdate(_ state: NWConnection.State) async {
SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug)
switch state { switch state {
case .ready: case .ready:
await self.readyState.markReady() self.markReady()
case .failed(let error): case .failed(let error):
await self.readyState.markFailed(error) self.markFailed(error)
self.emitEvent(.failed(error)) self.emitEvent(.failed(error))
case .cancelled: case .cancelled:
await self.readyState.markCancelled() self.markCancelled()
self.emitEvent(.cancelled) self.emitEvent(.cancelled)
case .setup, .preparing: case .setup, .preparing:
await self.readyState.markConnecting() self.markConnecting()
default: default:
() ()
} }
} }
func waitReady(timeout: Duration = .seconds(5)) async throws { func waitReady(timeout: Duration = .seconds(5)) async throws {
try await self.readyState.waitReady(timeout: timeout) try await self.waitReadyUntilStateChanged(timeout: timeout)
} }
func run() async -> SDLQUICClientExit { func run() async -> SDLQUICClientExit {
@ -130,16 +141,16 @@ actor SDLQUICClient {
group.addTask { group.addTask {
await self.readLoop() await self.readLoop()
} }
group.addTask { group.addTask {
await self.heartbeatLoop() await self.heartbeatLoop()
} }
let exit = await group.next() ?? .normal let exit = await group.next() ?? .normal
group.cancelAll() group.cancelAll()
await self.stop() await self.stop()
self.finishStreams() self.finishStreams()
return exit return exit
} }
} onCancel: { } onCancel: {
@ -148,14 +159,14 @@ actor SDLQUICClient {
} }
} }
} }
func send(type: SDLPacketType, data: Data) { func send(type: SDLPacketType, data: Data) {
var len = UInt16(data.count + 1).bigEndian var len = UInt16(data.count + 1).bigEndian
var packet = Data(Data(bytes: &len, count: 2)) var packet = Data(Data(bytes: &len, count: 2))
packet.append(type.rawValue) packet.append(type.rawValue)
packet.append(data) packet.append(data)
connection.send(content: packet, completion: .contentProcessed { [weak self] error in connection.send(content: packet, completion: .contentProcessed { [weak self] error in
if let error { if let error {
SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug) SDLLogger.log("[SDLQUICClient] send data get error: \(error)", for: .debug)
@ -165,36 +176,36 @@ actor SDLQUICClient {
} }
}) })
} }
private func heartbeatLoop() async -> SDLQUICClientExit { private func heartbeatLoop() async -> SDLQUICClientExit {
let timerStream = SDLAsyncTimerStream() let timerStream = SDLAsyncTimerStream()
timerStream.start(interval: .seconds(5)) timerStream.start(interval: .seconds(5))
for await _ in timerStream.stream { for await _ in timerStream.stream {
if Task.isCancelled { if Task.isCancelled {
break break
} }
self.send(type: .ping, data: Data()) self.send(type: .ping, data: Data())
} }
SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug) SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug)
return .cancelled return .cancelled
} }
func stop() async { func stop() async {
self.connection.cancel() self.connection.cancel()
await self.readyState.markCancelled() self.markCancelled()
self.finishStreams() self.finishStreams()
} }
private func emitEvent(_ event: SDLQUICEvent) { private func emitEvent(_ event: SDLQUICEvent) {
guard !self.didFinishStreams else { guard !self.didFinishStreams else {
return return
} }
self.eventCont.yield(event) self.eventCont.yield(event)
} }
private func finishStreams() { private func finishStreams() {
guard !self.didFinishStreams else { guard !self.didFinishStreams else {
return return
@ -204,115 +215,105 @@ actor SDLQUICClient {
self.messageCont.finish() self.messageCont.finish()
self.eventCont.finish() self.eventCont.finish()
} }
} }
// --MARK: Ready // --MARK: Ready
extension SDLQUICClient { extension SDLQUICClient {
actor SDLQUICReadyState {
enum State {
case idle
case connecting
case ready
case failed(Error)
case cancelled
}
private var state: State = .idle private func waitReadyUntilStateChanged(timeout: Duration) async throws {
private var continuations: [UUID: CheckedContinuation<Void, Error>] = [:] try Task.checkCancellation()
func waitReady(timeout: Duration) async throws { try await withTaskCancellationHandler {
let id = UUID() try await withCheckedThrowingContinuation { continuation in
let timeoutTask = Task { switch self.readyStatus {
try? await Task.sleep(for: timeout) case .ready:
if Task.isCancelled { continuation.resume()
return case .failed(let error):
} continuation.resume(throwing: error)
self.cancelWaiter(id: id, throwing: SDLQUICError.timeout) case .cancelled:
} continuation.resume(throwing: SDLQUICError.connectionCancelled)
case .idle, .connecting:
defer { guard self.readyContinuation == nil else {
timeoutTask.cancel() continuation.resume(throwing: SDLQUICError.waitReadyAlreadyInProgress)
} return
}
try await withTaskCancellationHandler { self.readyContinuation = continuation
try await withCheckedThrowingContinuation { continuation in self.readyTimeoutTask?.cancel()
self.addWaiter(id: id, continuation: continuation) self.readyTimeoutTask = Task {
} do {
} onCancel: { try await Task.sleep(for: timeout)
timeoutTask.cancel() if !Task.isCancelled {
Task { await self.cancelReadyWaiter(throwing: SDLQUICError.timeout)
await self.cancelWaiter(id: id, throwing: CancellationError()) }
} catch {
return
}
}
} }
} }
} } onCancel: {
Task {
private func addWaiter(id: UUID, continuation: CheckedContinuation<Void, Error>) { await self.cancelReadyWaiter(throwing: CancellationError())
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) { private func cancelReadyWaiter(throwing error: Error) {
guard let continuation = continuations.removeValue(forKey: id) else { guard let continuation = self.readyContinuation else {
return return
}
continuation.resume(throwing: error)
} }
func markConnecting() { self.readyContinuation = nil
switch state { self.readyTimeoutTask?.cancel()
case .idle: self.readyTimeoutTask = nil
state = .connecting continuation.resume(throwing: error)
default: }
break
} private func markConnecting() {
switch self.readyStatus {
case .idle:
self.readyStatus = .connecting
default:
break
}
}
private func markReady() {
self.readyStatus = .ready
self.resumeReadyWaiter()
}
private func markFailed(_ error: Error) {
self.readyStatus = .failed(error)
self.resumeReadyWaiter(throwing: error)
}
private func markCancelled() {
self.readyStatus = .cancelled
self.resumeReadyWaiter(throwing: SDLQUICError.connectionCancelled)
}
private func resumeReadyWaiter() {
guard let continuation = self.readyContinuation else {
return
} }
func markReady() { self.readyContinuation = nil
state = .ready self.readyTimeoutTask?.cancel()
self.readyTimeoutTask = nil
continuation.resume()
}
let list = continuations private func resumeReadyWaiter(throwing error: Error) {
continuations.removeAll() guard let continuation = self.readyContinuation else {
return
for continuation in list.values {
continuation.resume()
}
} }
func markFailed(_ error: Error) { self.readyContinuation = nil
state = .failed(error) self.readyTimeoutTask?.cancel()
self.readyTimeoutTask = nil
let list = continuations continuation.resume(throwing: error)
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())
}
}
} }
} }
@ -321,11 +322,11 @@ extension SDLQUICClient {
private func readLoop() async -> SDLQUICClientExit { private func readLoop() async -> SDLQUICClientExit {
var buffer = allocator.buffer(capacity: self.maxBufferSize) var buffer = allocator.buffer(capacity: self.maxBufferSize)
let threshold = self.maxBufferSize / 10 * 6 let threshold = self.maxBufferSize / 10 * 6
defer { defer {
self.messageCont.finish() self.messageCont.finish()
} }
do { do {
while !Task.isCancelled { while !Task.isCancelled {
let (isComplete, data) = try await self.readOnce() let (isComplete, data) = try await self.readOnce()
@ -335,14 +336,14 @@ extension SDLQUICClient {
if buffer.readerIndex > threshold { if buffer.readerIndex > threshold {
buffer.discardReadBytes() buffer.discardReadBytes()
} }
for frame in frames { for frame in frames {
if let message = decode(frame: frame) { if let message = decode(frame: frame) {
self.messageCont.yield(message) self.messageCont.yield(message)
} }
} }
} }
if isComplete { if isComplete {
return .transportClosed("receive complete") return .transportClosed("receive complete")
} }
@ -354,36 +355,36 @@ extension SDLQUICClient {
return .readFailed("\(error)") return .readFailed("\(error)")
} }
} }
// //
private func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] { private func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] {
guard buffer.readableBytes >= 2 else { guard buffer.readableBytes >= 2 else {
return [] return []
} }
var frames: [ByteBuffer] = [] var frames: [ByteBuffer] = []
while true { while true {
guard let len = buffer.getInteger(at: buffer.readerIndex, endianness: .big, as: UInt16.self) else { guard let len = buffer.getInteger(at: buffer.readerIndex, endianness: .big, as: UInt16.self) else {
break break
} }
if len > self.maxPacketSize { if len > self.maxPacketSize {
throw SDLQUICError.packetTooLarge throw SDLQUICError.packetTooLarge
} }
guard buffer.readableBytes >= len + 2 else { guard buffer.readableBytes >= len + 2 else {
break break
} }
buffer.moveReaderIndex(forwardBy: 2) buffer.moveReaderIndex(forwardBy: 2)
if let buf = buffer.readSlice(length: Int(len)) { if let buf = buffer.readSlice(length: Int(len)) {
frames.append(buf) frames.append(buf)
} }
} }
return frames return frames
} }
// //
private func readOnce() async throws -> (Bool, Data?) { private func readOnce() async throws -> (Bool, Data?) {
return try await withCheckedThrowingContinuation { cont in return try await withCheckedThrowingContinuation { cont in
@ -406,7 +407,7 @@ extension SDLQUICClient {
let packetType = SDLPacketType(rawValue: type) else { let packetType = SDLPacketType(rawValue: type) else {
return nil return nil
} }
switch packetType { switch packetType {
case .welcome: case .welcome:
guard let bytes = buffer.readBytes(length: buffer.readableBytes), guard let bytes = buffer.readBytes(length: buffer.readableBytes),
@ -414,7 +415,7 @@ extension SDLQUICClient {
return nil return nil
} }
return .welcome(welcome) return .welcome(welcome)
case .registerSuperAck: case .registerSuperAck:
guard let bytes = buffer.readBytes(length: buffer.readableBytes), guard let bytes = buffer.readBytes(length: buffer.readableBytes),
let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else { let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else {
@ -456,7 +457,7 @@ extension SDLQUICClient {
return .pong return .pong
default: default:
SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug) SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug)
return nil return nil
} }
} }
@ -464,50 +465,50 @@ extension SDLQUICClient {
// --MARK: quic // --MARK: quic
extension SDLQUICClient { extension SDLQUICClient {
enum QUICVerifier { enum QUICVerifier {
// Base64 // Base64
static let pinnedPublicKeyHashes = [ static let pinnedPublicKeyHashes = [
"Q41r6hbMWEVyxo6heNAH4Wx/TH5NNOWlNif9bewcJ3E=" "Q41r6hbMWEVyxo6heNAH4Wx/TH5NNOWlNif9bewcJ3E="
] ]
static func verify(trust: sec_trust_t, host: String) -> Bool { static func verify(trust: sec_trust_t, host: String) -> Bool {
let secTrust = sec_trust_copy_ref(trust).takeRetainedValue() let secTrust = sec_trust_copy_ref(trust).takeRetainedValue()
// --- Step 1: --- // --- Step 1: ---
var error: CFError? var error: CFError?
guard SecTrustEvaluateWithError(secTrust, &error) else { guard SecTrustEvaluateWithError(secTrust, &error) else {
SDLLogger.log("❌ 系统证书验证失败: \(error?.localizedDescription ?? "未知错误")", for: .debug) SDLLogger.log("❌ 系统证书验证失败: \(error?.localizedDescription ?? "未知错误")", for: .debug)
return false return false
} }
// --- Step 2: --- // --- Step 2: ---
let policy = SecPolicyCreateSSL(true, host as CFString) let policy = SecPolicyCreateSSL(true, host as CFString)
SecTrustSetPolicies(secTrust, policy) SecTrustSetPolicies(secTrust, policy)
guard SecTrustEvaluateWithError(secTrust, &error) else { guard SecTrustEvaluateWithError(secTrust, &error) else {
SDLLogger.log("❌ 主机名校验失败: \(error?.localizedDescription ?? "未知错误")", for: .debug) SDLLogger.log("❌ 主机名校验失败: \(error?.localizedDescription ?? "未知错误")", for: .debug)
return false return false
} }
// --- Step 3: --- // --- Step 3: ---
guard let chain = SecTrustCopyCertificateChain(secTrust) as? [SecCertificate], guard let chain = SecTrustCopyCertificateChain(secTrust) as? [SecCertificate],
let leafCertificate = chain.first else { let leafCertificate = chain.first else {
SDLLogger.log("❌ 无法获取证书链或叶子证书", for: .debug) SDLLogger.log("❌ 无法获取证书链或叶子证书", for: .debug)
return false return false
} }
// --- Step 4: --- // --- Step 4: ---
guard let publicKey = SecCertificateCopyKey(leafCertificate), guard let publicKey = SecCertificateCopyKey(leafCertificate),
let publicKeyData = SecKeyCopyExternalRepresentation(publicKey, nil) as Data? else { let publicKeyData = SecKeyCopyExternalRepresentation(publicKey, nil) as Data? else {
SDLLogger.log("❌ 无法提取公钥", for: .debug) SDLLogger.log("❌ 无法提取公钥", for: .debug)
return false return false
} }
// --- Step 5: SHA256 --- // --- Step 5: SHA256 ---
let hash = SHA256.hash(data: publicKeyData) let hash = SHA256.hash(data: publicKeyData)
let hashBase64 = Data(hash).base64EncodedString() let hashBase64 = Data(hash).base64EncodedString()
if pinnedPublicKeyHashes.contains(hashBase64) { if pinnedPublicKeyHashes.contains(hashBase64) {
SDLLogger.log("✅ 公钥校验通过", for: .debug) SDLLogger.log("✅ 公钥校验通过", for: .debug)
return true return true
@ -516,6 +517,6 @@ extension SDLQUICClient {
return false return false
} }
} }
} }
} }