调整生命周期的管理

This commit is contained in:
anlicheng 2026-05-05 15:57:14 +08:00
parent 3225d8fa17
commit 3bec145f23
2 changed files with 119 additions and 89 deletions

View File

@ -178,60 +178,33 @@ actor SDLContextActor {
try await Task.sleep(for: .seconds(0.5)) try await Task.sleep(for: .seconds(0.5))
SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)")
try await withTaskCancellationHandler {
try await withThrowingTaskGroup { group in try await withThrowingTaskGroup { group in
defer { defer {
group.cancelAll() group.cancelAll()
} }
// Actor AsyncStream
let (readyStream, readyContinuation) = AsyncStream<Void>.makeStream()
group.addTask { group.addTask {
for await event in quicClient.eventStream { for try await message in quicClient.messageStream {
try Task.checkCancellation()
switch event {
case .ready:
readyContinuation.yield()
case .failed(let error):
throw error
case .cancelled:
throw SDLQUICEvent.cancelled
case .writeFailed(let error):
throw error
}
}
}
group.addTask {
//
var it = readyStream.makeAsyncIterator()
await it.next()
try Task.checkCancellation()
await withThrowingTaskGroup { workerGroup in
workerGroup.addTask {
for await message in quicClient.messageStream {
try Task.checkCancellation() try Task.checkCancellation()
await self.handleQUICMessage(message: message) await self.handleQUICMessage(message: message)
} }
} }
workerGroup.addTask { group.addTask {
let timerStream = SDLAsyncTimerStream() while true {
timerStream.start(interval: .seconds(5)) try await Task.sleep(for: .seconds(5))
for await _ in timerStream.stream {
try Task.checkCancellation() try Task.checkCancellation()
quicClient.send(type: .ping, data: Data()) quicClient.send(type: .ping, data: Data())
} }
SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug) SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug)
} }
}
}
try await group.next() try await group.next()
} }
} onCancel: {
quicClient.stop()
}
} }
private func handleQUICMessage(message: SDLQUICInboundMessage) async { private func handleQUICMessage(message: SDLQUICInboundMessage) async {

View File

@ -15,55 +15,64 @@ import Security
enum SDLQUICError: Error { enum SDLQUICError: Error {
case connectionFailed(Error) case connectionFailed(Error)
case connectionCancelled case connectionCancelled
case writeFailed(Error)
case internalError(Error)
case timeout case timeout
case decodeError(String) case decodeError(String)
case packetTooLarge case packetTooLarge
case dataStreamClosed case dataStreamClosed
} }
enum SDLQUICEvent: Error {
case ready
case failed(Error)
case cancelled
case writeFailed(Error)
}
final class SDLQUICClient { final class SDLQUICClient {
enum State {
case idle
case running
case stopped
}
private var state: State = .idle
private let frameParser: SDLQUICFrameParser private let frameParser: SDLQUICFrameParser
// //
public var messageStream: AsyncStream<SDLQUICInboundMessage> public var messageStream: AsyncThrowingStream<SDLQUICInboundMessage, Error>
private let messageCont: AsyncStream<SDLQUICInboundMessage>.Continuation private let messageCont: AsyncThrowingStream<SDLQUICInboundMessage, Error>.Continuation
private var isMessageContinuationFinished: Bool = false
private var readTask: Task<Void, Never>? private var readTask: Task<Void, Never>?
// private var connection: NWConnection?
public var eventStream: AsyncStream<SDLQUICEvent>
private let eventCont: AsyncStream<SDLQUICEvent>.Continuation
private var isFinished: Bool = false
private let connection: NWConnection
private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 线 private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 线
private let queueKey = DispatchSpecificKey<Void>()
private let host: String
private let port: UInt16
init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) { init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) {
let options = NWProtocolTLS.Options() self.host = host
self.port = port
self.frameParser = SDLQUICFrameParser(maxBufferSize: maxBufferSize)
(self.messageStream, self.messageCont) = AsyncThrowingStream.makeStream(of: SDLQUICInboundMessage.self)
self.queue.setSpecific(key: self.queueKey, value: ())
}
func start() {
let options = NWProtocolTLS.Options()
sec_protocol_options_add_tls_application_protocol( sec_protocol_options_add_tls_application_protocol(
options.securityProtocolOptions, options.securityProtocolOptions,
"punchnet/1.0" "punchnet/1.0"
) )
self.frameParser = SDLQUICFrameParser(maxBufferSize: maxBufferSize)
(self.eventStream, self.eventCont) = AsyncStream.makeStream(of: SDLQUICEvent.self)
(self.messageStream, self.messageCont) = AsyncStream.makeStream(of: SDLQUICInboundMessage.self)
// //
sec_protocol_options_set_verify_block( sec_protocol_options_set_verify_block(
options.securityProtocolOptions, options.securityProtocolOptions,
{ metadata, trust, complete in { metadata, trust, complete in
// //
complete(TLSVerifier.verify(trust: trust, host: host)) complete(TLSVerifier.verify(trust: trust, host: self.host))
}, },
self.queue self.queue
) )
@ -73,49 +82,78 @@ final class SDLQUICClient {
// Network.framework // Network.framework
params.preferNoProxies = true params.preferNoProxies = true
self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) let connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
}
func start() {
connection.stateUpdateHandler = { [weak self] state in connection.stateUpdateHandler = { [weak self] state in
SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug)
switch state { switch state {
case .ready: case .ready:
self?.startReadTask() self?.startReadTask()
self?.eventCont.yield(.ready) self?.state = .running
case .failed(let error): case .failed(let error):
self?.eventCont.yield(.failed(error)) self?.finishMessageContinuationIfNeed(throwing: .connectionFailed(error))
case .cancelled: case .cancelled:
self?.eventCont.yield(.cancelled) self?.finishMessageContinuationIfNeed(throwing: .connectionCancelled)
default: default:
() ()
} }
} }
connection.start(queue: self.queue) connection.start(queue: self.queue)
self.connection = connection
}
private func finishMessageContinuationIfNeed(throwing error: SDLQUICError?) {
if DispatchQueue.getSpecific(key: queueKey) != nil {
self.finishMessageContinuationIfNeedOnQueue(throwing: error)
} else {
queue.async { [weak self] in
self?.finishMessageContinuationIfNeedOnQueue(throwing: error)
}
}
}
private func finishMessageContinuationIfNeedOnQueue(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() { private func startReadTask() {
self.readTask?.cancel() self.readTask?.cancel()
self.readTask = Task { self.readTask = Task {
do { do {
while !Task.isCancelled { while true {
try Task.checkCancellation()
let data = try await self.readOnce() let data = try await self.readOnce()
let frames = try self.frameParser.parseFrames(data: data) let frames = try self.frameParser.parseFrames(data: data)
for frame in frames { for frame in frames {
if let message = SDLQUICCodec.decode(frame: frame) { if let message = SDLQUICCodec.decode(frame: frame) {
self.messageCont.yield(message) self.messageCont.yield(message)
} else {
self.finishMessageContinuationIfNeed(throwing: .decodeError("invalid message"))
} }
} }
} }
} catch { } catch let err {
self.messageCont.finish() self.finishMessageContinuationIfNeed(throwing: .internalError(err))
} }
} }
} }
func send(type: SDLPacketType, data: Data) { func send(type: SDLPacketType, data: Data) {
var len = UInt16(data.count + 1).bigEndian 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)) var packet = Data(Data(bytes: &len, count: 2))
packet.append(type.rawValue) packet.append(type.rawValue)
packet.append(data) packet.append(data)
@ -123,14 +161,18 @@ final class SDLQUICClient {
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)
self?.eventCont.yield(.writeFailed(error)) self?.finishMessageContinuationIfNeed(throwing: .writeFailed(error))
} }
}) })
} }
private func readOnce() async throws -> Data { private func readOnce() async throws -> Data {
guard let connection = self.connection else {
throw SDLQUICError.connectionCancelled
}
return try await withCheckedThrowingContinuation { cont in return try await withCheckedThrowingContinuation { cont in
self.connection.receive(minimumIncompleteLength: 1, maximumLength: 64 * 1024) { data, _, isComplete, error in connection.receive(minimumIncompleteLength: 1, maximumLength: 64 * 1024) { data, _, isComplete, error in
if let error { if let error {
cont.resume(throwing: error) cont.resume(throwing: error)
return return
@ -146,10 +188,25 @@ final class SDLQUICClient {
} }
func stop() { func stop() {
if DispatchQueue.getSpecific(key: queueKey) != nil {
self.stopOnQueue()
} else {
queue.sync {
self.stopOnQueue()
}
}
}
private func stopOnQueue() {
guard self.state != .stopped else {
return
}
self.state = .stopped
self.readTask?.cancel() self.readTask?.cancel()
self.connection.cancel() self.connection?.cancel()
self.eventCont.finish() self.finishMessageContinuationIfNeed(throwing: nil)
self.messageCont.finish()
} }
} }