Files
2026-03-05 14:42:33 +03:00

425 lines
16 KiB
Python

"""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