diff --git a/Tun/Punchnet/Actors/SDLContextActor.swift b/Tun/Punchnet/Actors/SDLContextActor.swift index a0d40f6..cad17c9 100644 --- a/Tun/Punchnet/Actors/SDLContextActor.swift +++ b/Tun/Punchnet/Actors/SDLContextActor.swift @@ -267,7 +267,6 @@ actor SDLContextActor { quicClient.stop() } - try await quicClient.waitReady(timeout: .seconds(3)) // 这里必须等待quic的协商完成 try await Task.sleep(for: .seconds(0.3)) SDLLogger.log("[SDLContext] start quic client: \(self.config.serverHost)") @@ -278,7 +277,7 @@ actor SDLContextActor { } group.addTask { - for try await message in quicClient.messageStream() { + for await message in quicClient.messageStream { await self.handleQUICMessage(message: message) } } diff --git a/Tun/Punchnet/Actors/SDLQuicClient.swift b/Tun/Punchnet/Actors/SDLQuicClient.swift index 92e30dc..78cc4e8 100644 --- a/Tun/Punchnet/Actors/SDLQuicClient.swift +++ b/Tun/Punchnet/Actors/SDLQuicClient.swift @@ -18,6 +18,7 @@ enum SDLQUICError: Error { case timeout case decodeError(String) case packetTooLarge + case dataStreamClosed } enum SDLQUICEvent: Error { @@ -27,26 +28,28 @@ enum SDLQUICEvent: Error { } final class SDLQUICClient { - private let allocator = ByteBufferAllocator() - // 最大缓冲区区为2M - private let maxBufferSize: Int + private let frameParser: SDLQUICFrameParser - private static let maxPacketSize: Int = 64 * 1024 - - private let readyLatch = AsyncOneShot() + // 数据流 + public var messageStream: AsyncStream + private let messageCont: AsyncStream.Continuation + private var readTask: Task? // 事件流 public var eventStream: AsyncStream private let eventCont: AsyncStream.Continuation + + private var isFinished: Bool = false private let connection: NWConnection private let queue = DispatchQueue(label: "com.sdl.QUICClient.queue") // 专用队列保证线程安全 init(host: String, port: UInt16, maxBufferSize: Int = 2 * 1024 * 1024) { let options = NWProtocolQUIC.Options(alpn: ["punchnet/1.0"]) - - self.maxBufferSize = maxBufferSize + 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( @@ -61,36 +64,43 @@ final class SDLQUICClient { let params = NWParameters(quic: options) self.connection = NWConnection(host: .init(host), port: .init(rawValue: port)!, using: params) } - + func start() { connection.stateUpdateHandler = { [weak self] state in - Task { - await self?.handleConnectionStateUpdate(state) + SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) + switch state { + case .ready: + self?.startReadTask() + case .failed(let error): + self?.eventCont.yield(.failed(error)) + case .cancelled: + self?.eventCont.yield(.cancelled) + default: + () } } connection.start(queue: self.queue) } - private func handleConnectionStateUpdate(_ state: NWConnection.State) async { - SDLLogger.log("[SDLQUICClient] new state: \(state)", for: .debug) - switch state { - case .ready: - await self.readyLatch.succeed(()) - case .failed(let error): - await self.readyLatch.fail(error) - self.eventCont.yield(.failed(error)) - case .cancelled: - await self.readyLatch.fail(SDLQUICError.connectionCancelled) - self.eventCont.yield(.cancelled) - default: - () + private func startReadTask() { + self.readTask?.cancel() + self.readTask = Task { + do { + while !Task.isCancelled { + 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) + } + } + } + } catch { + self.messageCont.finish() + } } } - - func waitReady(timeout: Duration = .seconds(5)) async throws { - try await self.readyLatch.wait(timeout: timeout, timeoutError: SDLQUICError.timeout) - } - + func send(type: SDLPacketType, data: Data) { var len = UInt16(data.count + 1).bigEndian @@ -105,89 +115,84 @@ final class SDLQUICClient { } }) } + + private func readOnce() async throws -> Data { + return try await withCheckedThrowingContinuation { cont in + self.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() { + self.readTask?.cancel() self.connection.cancel() - Task { - await self.readyLatch.fail(SDLQUICError.connectionCancelled) - } self.eventCont.finish() + self.messageCont.finish() } } -// --MARK: Reader +// --MARK: 数据累加器 extension SDLQUICClient { - func messageStream() -> AsyncThrowingStream { - return AsyncThrowingStream { continuation in - var buffer = allocator.buffer(capacity: self.maxBufferSize) - let threshold = self.maxBufferSize / 10 * 6 + 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) - func readOnce() { - self.connection.receive(minimumIncompleteLength: 1, maximumLength: Self.maxPacketSize) { data, _, isComplete, error in - if let error { - continuation.finish(throwing: error) - return - } - - do { - if let data, !data.isEmpty { - buffer.writeBytes(data) - let frames = try Self.parseFrames(buffer: &buffer) - if buffer.readerIndex > threshold { - buffer.discardReadBytes() - } - - for frame in frames { - if let message = Self.decode(frame: frame) { - continuation.yield(message) - } - } - } - } catch let err { - continuation.finish(throwing: err) - return - } - - if isComplete { - continuation.finish() - } else { - readOnce() - } + 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) } } - readOnce() - } - } - - // 尝试解析数据 - private static func parseFrames(buffer: inout ByteBuffer) throws -> [ByteBuffer] { - guard buffer.readableBytes >= 2 else { - return [] + + let threshold = maxBufferSize / 10 * 6 + if buffer.readerIndex > threshold { + buffer.discardReadBytes() + } + + return frames } - 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) - } - } - - return frames } } @@ -195,64 +200,66 @@ extension SDLQUICClient { // --MARK: 编解码器 extension SDLQUICClient { - private 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 { + 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 } - return .welcome(welcome) - - case .registerSuperAck: - guard let bytes = buffer.readBytes(length: buffer.readableBytes), - let registerSuperAck = try? SDLRegisterSuperAck(serializedBytes: bytes) else { + + 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("SDLQUICClient decode Event Error", for: .debug) + return nil + } + return .event(event) + case .pong: + return .pong + default: + SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug) + 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("SDLQUICClient decode Event Error", for: .debug) - return nil - } - return .event(event) - case .pong: - return .pong - default: - SDLLogger.log("SDLQUICClient decode miss type: \(type)", for: .debug) - - return nil } } }