init commit
This commit is contained in:
@@ -0,0 +1,424 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user