"""Session orchestration for direct-USB LibreVNA protocol communication.""" from __future__ import annotations from collections import defaultdict, deque import logging import threading import time from typing import Callable, Protocol, TypeVar from .enums import HardwareFamily, PacketType, SyncMode from .exceptions import ( DeviceDisconnectedError, NackError, ParseError, ProtocolVersionMismatch, TimeoutError, ) from .models import DeviceInfo, DeviceStatus, Packet, StreamHandle, USBDeviceDescriptor from .protocol import ( FrameScanner, decode_packet_payload, encode_frame, encode_packet_payload, ensure_no_payload_types, parse_device_status, ) from .transport import USBTransport logger = logging.getLogger(__name__) class _IndexedPayload(Protocol): """Protocol for payload types carrying an integer `point_number` field.""" point_number: int TPacketPayload = TypeVar("TPacketPayload", bound=_IndexedPayload) class LibreVNASession: """Owns USB transport, packet queues, and request/response synchronization.""" def __init__(self) -> None: """Initialize transport, queues, synchronization primitives, and defaults.""" self._scanner = FrameScanner() self._transport = USBTransport( on_data=self._on_transport_data, on_disconnect=self._on_transport_disconnect, ) self._incoming: dict[PacketType, deque[Packet]] = defaultdict(deque) self._subscribers: dict[PacketType, set[Callable[[Packet], None]]] = defaultdict(set) self._incoming_cv = threading.Condition() self._ack_cv = threading.Condition() self._send_lock = threading.Lock() self._awaiting_ack = False self._ack_result: PacketType | None = None self._fatal_error: Exception | None = None self._device_info: DeviceInfo | None = None self._device_status: DeviceStatus | None = None self._hardware_family: HardwareFamily = HardwareFamily.UNKNOWN self._default_sync_mode = SyncMode.DISABLED self._default_sync_master = False @property def is_connected(self) -> bool: """Return `True` when USB connection is active.""" return self._transport.is_connected @property def connected_serial(self) -> str | None: """Serial number of currently connected device, when available.""" return self._transport.connected_serial @property def hardware_family(self) -> HardwareFamily: """Hardware family inferred from `DeviceInfo` packet.""" return self._hardware_family @property def default_sync_mode(self) -> SyncMode: """Default sync mode used by controllers when settings do not override.""" return self._default_sync_mode @property def default_sync_master(self) -> bool: """Default sync master flag used by controllers.""" return self._default_sync_master def set_default_sync(self, mode: SyncMode, master: bool) -> None: """Set session-level default synchronization settings.""" self._default_sync_mode = mode self._default_sync_master = master @staticmethod def list_devices() -> list[USBDeviceDescriptor]: """Enumerate available direct-USB LibreVNA devices.""" return USBTransport.list_devices() def connect( self, serial: str | None = None, *, strict_protocol_version: int = 14, timeout_s: float = 1.0, ) -> None: """Connect transport and perform startup handshake.""" logger.info( "Opening USB transport (serial=%s, strict_protocol=%d, timeout=%.2fs)", serial, strict_protocol_version, timeout_s, ) self._fatal_error = None self._scanner.clear() self._incoming.clear() self._transport.connect(serial=serial, timeout_s=timeout_s) try: info_packet = self.request( PacketType.REQUEST_DEVICE_INFO, PacketType.DEVICE_INFO, timeout_s=timeout_s, ) if not isinstance(info_packet.payload, DeviceInfo): raise ParseError("Decoded DeviceInfo payload has unexpected type") self._device_info = info_packet.payload self._hardware_family = self._device_info.hardware_family logger.info( "Handshake OK: fw=%s protocol=%d family=%s ports=%d", self._device_info.firmware_version, self._device_info.protocol_version, self._device_info.hardware_family.name, self._device_info.num_ports, ) if self._device_info.protocol_version != strict_protocol_version: raise ProtocolVersionMismatch( "Protocol version mismatch: " f"device={self._device_info.protocol_version}, " f"required={strict_protocol_version}" ) self.get_device_status(timeout_s=timeout_s) logger.debug("Initial device status request completed") except Exception: logger.exception("Connect handshake failed, closing session") self.disconnect() raise def disconnect(self) -> None: """Disconnect USB transport and clear session state.""" logger.debug("Closing USB transport and clearing session queues") self._transport.disconnect() with self._incoming_cv: self._incoming.clear() self._subscribers.clear() self._incoming_cv.notify_all() with self._ack_cv: self._awaiting_ack = False self._ack_result = None self._ack_cv.notify_all() def send(self, packet: Packet, *, require_ack: bool = True, timeout_s: float = 0.5) -> None: """Send one protocol packet and optionally wait for ACK/NACK.""" payload = encode_packet_payload(packet.type, packet.payload) ensure_no_payload_types(packet.type, payload) frame = encode_frame(Packet(type=packet.type, payload=payload)) logger.debug( "TX packet=%s payload=%dB require_ack=%s timeout=%.2fs", packet.type.name, len(payload), require_ack, timeout_s, ) with self._send_lock: if require_ack: with self._ack_cv: self._awaiting_ack = True self._ack_result = None self._transport.write(frame, timeout_s=timeout_s) if require_ack: deadline = time.monotonic() + timeout_s with self._ack_cv: while self._ack_result is None: self._raise_if_fatal_locked() remaining = deadline - time.monotonic() if remaining <= 0: self._awaiting_ack = False raise TimeoutError(f"Timeout waiting for ACK for packet {packet.type.name}") self._ack_cv.wait(timeout=remaining) result = self._ack_result self._awaiting_ack = False self._ack_result = None if result == PacketType.NACK: raise NackError(f"Received NACK for packet {packet.type.name}") logger.debug("ACK received for packet %s", packet.type.name) def request( self, packet_type: PacketType, response_type: PacketType, *, timeout_s: float = 1.0, ) -> Packet: """Send no-payload request packet and wait for one response packet type.""" logger.debug( "Request start: packet=%s expect=%s timeout=%.2fs", packet_type.name, response_type.name, timeout_s, ) self.clear_queue(response_type) self.send(Packet(type=packet_type), require_ack=True, timeout_s=timeout_s) packet = self.wait_for_packet(response_type, timeout_s=timeout_s) logger.debug("Request complete: received %s", response_type.name) return packet def collect_indexed_payloads( self, *, packet_type: PacketType, expected_points: int, timeout_s: float, payload_type: type[TPacketPayload], payload_error: str, ) -> list[TPacketPayload]: """Collect indexed payloads by `point_number` into ascending order. Missing points are omitted; callers decide whether this is acceptable. """ if expected_points <= 0: raise ValueError("expected_points must be > 0") deadline = time.monotonic() + timeout_s collected: list[TPacketPayload | None] = [None] * expected_points received = 0 while received < expected_points: remaining = deadline - time.monotonic() if remaining <= 0: break packet = self.wait_for_packet(packet_type, timeout_s=remaining) payload = packet.payload if not isinstance(payload, payload_type): raise ParseError(payload_error) point_number = payload.point_number if point_number < 0 or point_number >= expected_points: raise ParseError( f"Received out-of-range point index {point_number}, expected 0..{expected_points - 1}" ) if collected[point_number] is None: received += 1 collected[point_number] = payload return [item for item in collected if item is not None] def wait_for_packet(self, packet_type: PacketType, *, timeout_s: float = 1.0) -> Packet: """Wait for the next packet of requested type.""" deadline = time.monotonic() + timeout_s with self._incoming_cv: while True: self._raise_if_fatal_locked() queue = self._incoming[packet_type] if queue: logger.debug( "Dequeued packet %s (remaining=%d)", packet_type.name, len(queue) - 1, ) return queue.popleft() remaining = deadline - time.monotonic() if remaining <= 0: raise TimeoutError(f"Timeout waiting for packet {packet_type.name}") self._incoming_cv.wait(timeout=remaining) def clear_queue(self, packet_type: PacketType) -> None: """Drop queued packets of a given type.""" with self._incoming_cv: dropped = len(self._incoming[packet_type]) self._incoming[packet_type].clear() if dropped: logger.debug("Cleared %d queued packet(s) of type %s", dropped, packet_type.name) def get_device_info(self) -> DeviceInfo: """Return cached `DeviceInfo` from successful `connect()` handshake.""" if self._device_info is None: raise DeviceDisconnectedError("Device info is not available before connect()") return self._device_info def get_device_status(self, *, timeout_s: float = 1.0) -> DeviceStatus: """Request and return current `DeviceStatus`.""" packet = self.request( PacketType.REQUEST_DEVICE_STATUS, PacketType.DEVICE_STATUS, timeout_s=timeout_s, ) if not isinstance(packet.payload, (bytes, bytearray, memoryview)): raise ParseError("DeviceStatus packet payload has unexpected type") status = parse_device_status(bytes(packet.payload), self._hardware_family) self._device_status = status logger.debug( "Device status: source_locked=%s lo_locked=%s adc_overload=%s unlevel=%s", status.source_locked, status.lo_locked, status.adc_overload, status.unlevel, ) return status def subscribe(self, packet_type: PacketType, callback: Callable[[Packet], None]) -> StreamHandle: """Register packet callback and return handle for unsubscription.""" stop_event = threading.Event() with self._incoming_cv: self._subscribers[packet_type].add(callback) logger.debug( "Subscriber added for %s (count=%d)", packet_type.name, len(self._subscribers[packet_type]), ) def _close() -> None: """Unsubscribe callback from session packet subscribers.""" with self._incoming_cv: callbacks = self._subscribers.get(packet_type) if callbacks is not None: callbacks.discard(callback) logger.debug( "Subscriber removed for %s (count=%d)", packet_type.name, len(callbacks), ) return StreamHandle(stop_event=stop_event, close_callback=_close) def _on_transport_data(self, chunk: bytes) -> None: """Decode incoming USB chunk and route packets to queues/subscribers.""" try: packets = self._scanner.feed(chunk) if logger.isEnabledFor(logging.DEBUG) and packets: logger.debug("RX chunk=%dB decoded_packets=%d", len(chunk), len(packets)) except Exception as exc: self._set_fatal_error(exc) return for packet in packets: try: decoded_payload = decode_packet_payload(packet) except Exception as exc: self._set_fatal_error(exc) return decoded_packet = Packet(type=packet.type, payload=decoded_payload) self._dispatch_packet(decoded_packet) def _on_transport_disconnect(self, exc: Exception) -> None: """Receive asynchronous transport disconnect event.""" self._set_fatal_error(exc) def _dispatch_packet(self, packet: Packet) -> None: """Route packet into ack waiter, queue, and subscriber callbacks.""" if logger.isEnabledFor(logging.DEBUG) and packet.type not in { PacketType.VNA_DATAPOINT, }: logger.debug("Dispatch packet %s", packet.type.name) if packet.type in {PacketType.ACK, PacketType.NACK}: with self._ack_cv: if self._awaiting_ack and self._ack_result is None: self._ack_result = packet.type self._ack_cv.notify_all() return callbacks: list[Callable[[Packet], None]] = [] with self._incoming_cv: self._incoming[packet.type].append(packet) callbacks = list(self._subscribers.get(packet.type, set())) self._incoming_cv.notify_all() for callback in callbacks: try: callback(packet) except Exception as exc: logger.exception("Subscriber callback failed for packet %s", packet.type.name) self._set_fatal_error(exc) return def _set_fatal_error(self, exc: Exception) -> None: """Mark session as failed and wake all waiting operations.""" with self._incoming_cv: with self._ack_cv: if self._fatal_error is None: self._fatal_error = exc logger.error("Session fatal error: %s", exc, exc_info=exc) self._incoming_cv.notify_all() self._ack_cv.notify_all() def _raise_if_fatal_locked(self) -> None: """Raise stored fatal error (called while condition lock is held).""" if self._fatal_error is None: return exc = self._fatal_error if isinstance(exc, DeviceDisconnectedError): raise exc raise DeviceDisconnectedError(f"Session stopped due to transport/protocol error: {exc}") from exc