diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index b134614..9385255 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -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.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() } } diff --git a/Tun/Punchnet/Actors/SDLQuicClient.swift b/Tun/Punchnet/Actors/SDLQuicClient.swift index 99bc7a0..fcb0bd7 100644 --- a/Tun/Punchnet/Actors/SDLQuicClient.swift +++ b/Tun/Punchnet/Actors/SDLQuicClient.swift @@ -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 - private let messageCont: AsyncStream.Continuation + public var messageStream: AsyncThrowingStream + private let messageCont: AsyncThrowingStream.Continuation + private var isMessageContinuationFinished: Bool = false + private var readTask: Task? - // 事件流 - public var eventStream: AsyncStream - private let eventCont: AsyncStream.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() + + 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) } }