调整生命周期的管理
This commit is contained in:
parent
3225d8fa17
commit
3bec145f23
@ -178,59 +178,32 @@ actor SDLContextActor {
|
||||
try await Task.sleep(for: .seconds(0.5))
|
||||
SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)")
|
||||
|
||||
try await withThrowingTaskGroup { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
// 创建一个简单的异步状态等待机制(可以用一个 Actor 或者 AsyncStream 模拟)
|
||||
let (readyStream, readyContinuation) = AsyncStream<Void>.makeStream()
|
||||
|
||||
group.addTask {
|
||||
for await event in quicClient.eventStream {
|
||||
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
|
||||
try await withTaskCancellationHandler {
|
||||
try await withThrowingTaskGroup { group in
|
||||
defer {
|
||||
group.cancelAll()
|
||||
}
|
||||
|
||||
group.addTask {
|
||||
for try await message in quicClient.messageStream {
|
||||
try Task.checkCancellation()
|
||||
await self.handleQUICMessage(message: message)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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()
|
||||
await self.handleQUICMessage(message: message)
|
||||
}
|
||||
}
|
||||
|
||||
workerGroup.addTask {
|
||||
let timerStream = SDLAsyncTimerStream()
|
||||
timerStream.start(interval: .seconds(5))
|
||||
|
||||
for await _ in timerStream.stream {
|
||||
try Task.checkCancellation()
|
||||
quicClient.send(type: .ping, data: Data())
|
||||
}
|
||||
SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug)
|
||||
group.addTask {
|
||||
while true {
|
||||
try await Task.sleep(for: .seconds(5))
|
||||
try Task.checkCancellation()
|
||||
quicClient.send(type: .ping, data: Data())
|
||||
}
|
||||
SDLLogger.log("[SDLQUICClient] udp pingTask cancel", for: .debug)
|
||||
}
|
||||
|
||||
try await group.next()
|
||||
}
|
||||
|
||||
try await group.next()
|
||||
} onCancel: {
|
||||
quicClient.stop()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -15,55 +15,64 @@ 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
|
||||
}
|
||||
|
||||
enum SDLQUICEvent: Error {
|
||||
case ready
|
||||
case failed(Error)
|
||||
case cancelled
|
||||
case writeFailed(Error)
|
||||
}
|
||||
|
||||
final class SDLQUICClient {
|
||||
enum State {
|
||||
case idle
|
||||
case running
|
||||
case stopped
|
||||
}
|
||||
|
||||
private var state: State = .idle
|
||||
|
||||
private let frameParser: SDLQUICFrameParser
|
||||
|
||||
// 数据流
|
||||
public var messageStream: AsyncStream<SDLQUICInboundMessage>
|
||||
private let messageCont: AsyncStream<SDLQUICInboundMessage>.Continuation
|
||||
public var messageStream: AsyncThrowingStream<SDLQUICInboundMessage, Error>
|
||||
private let messageCont: AsyncThrowingStream<SDLQUICInboundMessage, Error>.Continuation
|
||||
private var isMessageContinuationFinished: Bool = false
|
||||
|
||||
private var readTask: Task<Void, Never>?
|
||||
|
||||
// 事件流
|
||||
public var eventStream: AsyncStream<SDLQUICEvent>
|
||||
private let eventCont: AsyncStream<SDLQUICEvent>.Continuation
|
||||
|
||||
private var isFinished: Bool = false
|
||||
|
||||
private let connection: NWConnection
|
||||
private var connection: NWConnection?
|
||||
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) {
|
||||
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(
|
||||
options.securityProtocolOptions,
|
||||
"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(
|
||||
options.securityProtocolOptions,
|
||||
{ metadata, trust, complete in
|
||||
// 执行公钥校验
|
||||
complete(TLSVerifier.verify(trust: trust, host: host))
|
||||
complete(TLSVerifier.verify(trust: trust, host: self.host))
|
||||
},
|
||||
self.queue
|
||||
)
|
||||
@ -73,64 +82,97 @@ final class SDLQUICClient {
|
||||
// 关键:让 Network.framework 忽略系统代理
|
||||
params.preferNoProxies = true
|
||||
|
||||
self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
|
||||
}
|
||||
|
||||
func start() {
|
||||
let connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params)
|
||||
|
||||
connection.stateUpdateHandler = { [weak self] state in
|
||||
SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug)
|
||||
switch state {
|
||||
case .ready:
|
||||
self?.startReadTask()
|
||||
self?.eventCont.yield(.ready)
|
||||
self?.state = .running
|
||||
case .failed(let error):
|
||||
self?.eventCont.yield(.failed(error))
|
||||
self?.finishMessageContinuationIfNeed(throwing: .connectionFailed(error))
|
||||
case .cancelled:
|
||||
self?.eventCont.yield(.cancelled)
|
||||
self?.finishMessageContinuationIfNeed(throwing: .connectionCancelled)
|
||||
default:
|
||||
()
|
||||
}
|
||||
}
|
||||
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() {
|
||||
self.readTask?.cancel()
|
||||
self.readTask = Task {
|
||||
do {
|
||||
while !Task.isCancelled {
|
||||
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 {
|
||||
self.messageCont.finish()
|
||||
} 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 {
|
||||
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 {
|
||||
guard let connection = self.connection else {
|
||||
throw SDLQUICError.connectionCancelled
|
||||
}
|
||||
|
||||
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 {
|
||||
cont.resume(throwing: error)
|
||||
return
|
||||
@ -144,12 +186,27 @@ final class SDLQUICClient {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
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.connection.cancel()
|
||||
self.eventCont.finish()
|
||||
self.messageCont.finish()
|
||||
self.connection?.cancel()
|
||||
self.finishMessageContinuationIfNeed(throwing: nil)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user