425 lines
16 KiB
Python
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
|