461 lines
16 KiB
Python
461 lines
16 KiB
Python
"""Event-driven locator TCP service fed by already-consumed GUI GPR results."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import contextlib
|
|
from dataclasses import dataclass
|
|
import json
|
|
import logging
|
|
import math
|
|
import queue
|
|
import struct
|
|
import threading
|
|
from typing import Any
|
|
|
|
from python_app.models.dataset_model import ResultCollection
|
|
from python_app.orchestration.gpr_locator import (
|
|
build_locator_payload,
|
|
locator_observations_from_collection,
|
|
)
|
|
|
|
_PACKET_HEADER_STRUCT = struct.Struct("<II")
|
|
|
|
|
|
def encode_packet(payload: dict[str, Any], device_id: int) -> bytes:
|
|
"""Serialize a JSON payload with the protocol binary header."""
|
|
payload_bytes = json.dumps(
|
|
payload,
|
|
ensure_ascii=True,
|
|
separators=(",", ":"),
|
|
).encode("utf-8")
|
|
return _PACKET_HEADER_STRUCT.pack(device_id, len(payload_bytes)) + payload_bytes
|
|
|
|
|
|
def decode_packet(header_bytes: bytes, payload_bytes: bytes) -> tuple[int, Any]:
|
|
"""Decode one protocol packet from its binary header and JSON payload."""
|
|
if len(header_bytes) != _PACKET_HEADER_STRUCT.size:
|
|
raise ValueError(f"Packet header must be exactly {_PACKET_HEADER_STRUCT.size} bytes long.")
|
|
|
|
device_id, payload_length = _PACKET_HEADER_STRUCT.unpack(header_bytes)
|
|
if payload_length != len(payload_bytes):
|
|
raise ValueError("Payload length does not match the header value.")
|
|
|
|
try:
|
|
payload = json.loads(payload_bytes.decode("utf-8"))
|
|
except UnicodeDecodeError as error:
|
|
raise ValueError("Payload is not valid UTF-8.") from error
|
|
except json.JSONDecodeError as error:
|
|
raise ValueError("Payload is not valid JSON.") from error
|
|
|
|
return device_id, payload
|
|
|
|
|
|
def parse_vlc(payload: dict[str, Any]) -> float:
|
|
"""Validate and normalize inbound speed payload."""
|
|
try:
|
|
vlc = float(payload["vlc"])
|
|
except (KeyError, TypeError, ValueError) as error:
|
|
raise ValueError("Payload field 'vlc' must be numeric.") from error
|
|
|
|
if not math.isfinite(vlc):
|
|
raise ValueError("Payload field 'vlc' must be finite.")
|
|
return vlc
|
|
|
|
|
|
def format_payload_for_log(payload: Any) -> str:
|
|
"""Return compact JSON-ish payload text for logs."""
|
|
return json.dumps(payload, ensure_ascii=True, separators=(",", ":"))
|
|
|
|
|
|
def decode_packet_for_log(packet: bytes) -> tuple[int, str]:
|
|
"""Decode an outbound packet into `(device_id, payload_text)` for logging."""
|
|
if len(packet) < _PACKET_HEADER_STRUCT.size:
|
|
raise ValueError("Packet is shorter than the locator header")
|
|
header_bytes = packet[: _PACKET_HEADER_STRUCT.size]
|
|
payload_bytes = packet[_PACKET_HEADER_STRUCT.size :]
|
|
device_id, payload = decode_packet(header_bytes, payload_bytes)
|
|
return device_id, format_payload_for_log(payload)
|
|
|
|
|
|
def format_peer_name(writer: asyncio.StreamWriter) -> str:
|
|
"""Return a readable peer address for logs."""
|
|
peer_name = writer.get_extra_info("peername")
|
|
if isinstance(peer_name, tuple) and len(peer_name) >= 2:
|
|
return f"{peer_name[0]}:{peer_name[1]}"
|
|
return str(peer_name or "unknown")
|
|
|
|
|
|
async def read_packet_with_limit(reader: asyncio.StreamReader, max_payload_bytes: int) -> tuple[int, Any]:
|
|
"""Read and decode a single packet using the requested payload limit."""
|
|
header_bytes = await reader.readexactly(_PACKET_HEADER_STRUCT.size)
|
|
_, payload_length = _PACKET_HEADER_STRUCT.unpack(header_bytes)
|
|
if payload_length > int(max_payload_bytes):
|
|
raise ValueError(
|
|
"Payload length %d exceeds the %d byte limit."
|
|
% (payload_length, int(max_payload_bytes))
|
|
)
|
|
payload_bytes = await reader.readexactly(payload_length)
|
|
return decode_packet(header_bytes, payload_bytes)
|
|
|
|
|
|
@dataclass(eq=False, slots=True)
|
|
class _ClientConnection:
|
|
"""Runtime state for one connected locator client."""
|
|
|
|
writer: asyncio.StreamWriter
|
|
peer_name: str
|
|
queue: asyncio.Queue[bytes]
|
|
closed: bool = False
|
|
|
|
|
|
class LocatorTcpService:
|
|
"""Background-thread TCP service for locator packets."""
|
|
|
|
def __init__(
|
|
self,
|
|
host: str,
|
|
port: int,
|
|
*,
|
|
device_id: int,
|
|
protocol_version: int,
|
|
max_payload_bytes: int,
|
|
client_queue_size: int,
|
|
logger_name: str,
|
|
logger: logging.Logger | None = None,
|
|
) -> None:
|
|
"""Create a stopped service instance."""
|
|
self._host = host
|
|
self._port = int(port)
|
|
self._device_id = int(device_id)
|
|
self._protocol_version = int(protocol_version)
|
|
self._max_payload_bytes = int(max_payload_bytes)
|
|
self._logger = logger or logging.getLogger(str(logger_name))
|
|
self._client_queue_size = int(client_queue_size)
|
|
self._speed_updates: queue.Queue[float] = queue.Queue()
|
|
self._log_updates: queue.Queue[str] = queue.Queue()
|
|
self._loop: asyncio.AbstractEventLoop | None = None
|
|
self._server: asyncio.AbstractServer | None = None
|
|
self._thread: threading.Thread | None = None
|
|
self._startup_event = threading.Event()
|
|
self._startup_error: Exception | None = None
|
|
self._clients: set[_ClientConnection] = set()
|
|
self._snapshot_lock = threading.Lock()
|
|
self._latest_packet: bytes | None = None
|
|
|
|
@property
|
|
def host(self) -> str:
|
|
"""Return bind host."""
|
|
return self._host
|
|
|
|
@property
|
|
def port(self) -> int:
|
|
"""Return bind port."""
|
|
return self._port
|
|
|
|
def start(self) -> None:
|
|
"""Start the background event loop and TCP listener."""
|
|
if self.is_running():
|
|
return
|
|
|
|
self._startup_event = threading.Event()
|
|
self._startup_error = None
|
|
self._thread = threading.Thread(
|
|
target=self._thread_main,
|
|
name="locator-tcp-service",
|
|
daemon=True,
|
|
)
|
|
self._thread.start()
|
|
|
|
if not self._startup_event.wait(timeout=5.0):
|
|
raise RuntimeError("Timed out waiting for locator TCP service startup.")
|
|
|
|
if self._startup_error is not None:
|
|
error = self._startup_error
|
|
self.stop()
|
|
raise RuntimeError(f"Failed to start locator TCP service: {error}") from error
|
|
|
|
def stop(self) -> None:
|
|
"""Stop listener, disconnect clients, and join the background thread."""
|
|
loop = self._loop
|
|
thread = self._thread
|
|
|
|
if loop is not None:
|
|
with contextlib.suppress(RuntimeError):
|
|
loop.call_soon_threadsafe(loop.stop)
|
|
|
|
if thread is not None:
|
|
thread.join(timeout=5.0)
|
|
|
|
self._thread = None
|
|
self._loop = None
|
|
self._server = None
|
|
self._clients.clear()
|
|
|
|
def is_running(self) -> bool:
|
|
"""Return whether the background loop is alive."""
|
|
return self._thread is not None and self._thread.is_alive() and self._loop is not None
|
|
|
|
def publish_collection(
|
|
self,
|
|
collection: ResultCollection,
|
|
min_score: float,
|
|
*,
|
|
visible_bounds: tuple[float, float, float, float] | None = None,
|
|
object_draw_limits: tuple[int, int] | None = None,
|
|
) -> None:
|
|
"""Publish one locator payload derived from a GPR result collection."""
|
|
observations = locator_observations_from_collection(
|
|
collection,
|
|
min_score,
|
|
visible_bounds=visible_bounds,
|
|
object_draw_limits=object_draw_limits,
|
|
)
|
|
payload = build_locator_payload(
|
|
observations,
|
|
protocol_version=self._protocol_version,
|
|
status=1,
|
|
)
|
|
self._publish_packet(encode_packet(payload, device_id=self._device_id))
|
|
|
|
def publish_empty(self) -> None:
|
|
"""Publish an empty locator snapshot."""
|
|
payload = build_locator_payload(
|
|
[],
|
|
protocol_version=self._protocol_version,
|
|
status=1,
|
|
)
|
|
self._publish_packet(encode_packet(payload, device_id=self._device_id))
|
|
|
|
def drain_speed_updates(self) -> float | None:
|
|
"""Drain queued speed updates and return the newest one, if any."""
|
|
latest: float | None = None
|
|
while True:
|
|
try:
|
|
latest = float(self._speed_updates.get_nowait())
|
|
except queue.Empty:
|
|
return latest
|
|
|
|
def drain_log_updates(self) -> list[str]:
|
|
"""Drain queued socket traffic log lines."""
|
|
lines: list[str] = []
|
|
while True:
|
|
try:
|
|
lines.append(str(self._log_updates.get_nowait()))
|
|
except queue.Empty:
|
|
return lines
|
|
|
|
def _queue_log_update(self, message: str) -> None:
|
|
"""Queue one socket traffic line for the GUI runtime log."""
|
|
self._log_updates.put(str(message))
|
|
|
|
def _log_socket_traffic(self, message: str) -> None:
|
|
"""Log socket traffic to both Python logging and the GUI-visible queue."""
|
|
self._logger.info(message)
|
|
self._queue_log_update(message)
|
|
|
|
def _publish_packet(self, packet: bytes) -> None:
|
|
"""Store latest packet and broadcast it to all connected clients."""
|
|
with self._snapshot_lock:
|
|
self._latest_packet = packet
|
|
|
|
loop = self._loop
|
|
if loop is None:
|
|
return
|
|
|
|
with contextlib.suppress(RuntimeError):
|
|
loop.call_soon_threadsafe(self._broadcast_packet, packet)
|
|
|
|
def _get_latest_packet(self) -> bytes | None:
|
|
"""Return the latest stored packet snapshot."""
|
|
with self._snapshot_lock:
|
|
return self._latest_packet
|
|
|
|
def _thread_main(self) -> None:
|
|
"""Own the event loop and TCP listener lifecycle."""
|
|
loop = asyncio.new_event_loop()
|
|
self._loop = loop
|
|
asyncio.set_event_loop(loop)
|
|
|
|
try:
|
|
self._server = loop.run_until_complete(
|
|
asyncio.start_server(self._handle_client, self._host, self._port)
|
|
)
|
|
except Exception as exc: # noqa: BLE001
|
|
self._startup_error = exc
|
|
self._startup_event.set()
|
|
self._loop = None
|
|
asyncio.set_event_loop(None)
|
|
loop.close()
|
|
return
|
|
|
|
self._startup_event.set()
|
|
try:
|
|
loop.run_forever()
|
|
finally:
|
|
with contextlib.suppress(Exception):
|
|
loop.run_until_complete(self._shutdown_async())
|
|
asyncio.set_event_loop(None)
|
|
loop.close()
|
|
self._server = None
|
|
self._loop = None
|
|
|
|
async def _shutdown_async(self) -> None:
|
|
"""Close listener and all active client connections."""
|
|
server = self._server
|
|
if server is not None:
|
|
server.close()
|
|
await server.wait_closed()
|
|
|
|
clients = list(self._clients)
|
|
self._clients.clear()
|
|
for client in clients:
|
|
client.closed = True
|
|
client.writer.close()
|
|
|
|
for client in clients:
|
|
with contextlib.suppress(BrokenPipeError, ConnectionResetError):
|
|
await client.writer.wait_closed()
|
|
|
|
pending = [
|
|
task
|
|
for task in asyncio.all_tasks()
|
|
if task is not asyncio.current_task()
|
|
]
|
|
for task in pending:
|
|
task.cancel()
|
|
for task in pending:
|
|
with contextlib.suppress(asyncio.CancelledError, Exception):
|
|
await task
|
|
|
|
async def _handle_client(
|
|
self,
|
|
reader: asyncio.StreamReader,
|
|
writer: asyncio.StreamWriter,
|
|
) -> None:
|
|
"""Handle one client until disconnect or protocol failure."""
|
|
peer_name = format_peer_name(writer)
|
|
client = _ClientConnection(
|
|
writer=writer,
|
|
peer_name=peer_name,
|
|
queue=asyncio.Queue(maxsize=self._client_queue_size),
|
|
)
|
|
self._clients.add(client)
|
|
self._logger.info("Locator client connected: %s", peer_name)
|
|
|
|
latest_packet = self._get_latest_packet()
|
|
if latest_packet is not None:
|
|
self._enqueue_packet(client, latest_packet)
|
|
|
|
send_task = asyncio.create_task(
|
|
self._send_packets(client),
|
|
name=f"locator_send:{peer_name}",
|
|
)
|
|
receive_task = asyncio.create_task(
|
|
self._receive_packets(reader, client),
|
|
name=f"locator_receive:{peer_name}",
|
|
)
|
|
|
|
done, pending = await asyncio.wait(
|
|
{send_task, receive_task},
|
|
return_when=asyncio.FIRST_COMPLETED,
|
|
)
|
|
|
|
for task in pending:
|
|
task.cancel()
|
|
for task in pending:
|
|
with contextlib.suppress(asyncio.CancelledError):
|
|
await task
|
|
|
|
self._clients.discard(client)
|
|
client.closed = True
|
|
writer.close()
|
|
with contextlib.suppress(BrokenPipeError, ConnectionResetError):
|
|
await writer.wait_closed()
|
|
|
|
for task in done:
|
|
exception = task.exception()
|
|
if exception is None:
|
|
continue
|
|
if isinstance(exception, asyncio.IncompleteReadError):
|
|
self._logger.info("Locator client closed the connection: %s", peer_name)
|
|
continue
|
|
if isinstance(exception, (BrokenPipeError, ConnectionResetError)):
|
|
self._logger.info("Locator connection lost: %s", peer_name)
|
|
continue
|
|
if isinstance(exception, ValueError):
|
|
self._logger.warning(
|
|
"Closing locator client %s after protocol error: %s",
|
|
peer_name,
|
|
exception,
|
|
)
|
|
continue
|
|
self._logger.error(
|
|
"Unexpected locator client error: %s",
|
|
peer_name,
|
|
exc_info=(type(exception), exception, exception.__traceback__),
|
|
)
|
|
|
|
self._logger.info("Locator client disconnected: %s", peer_name)
|
|
|
|
async def _send_packets(self, client: _ClientConnection) -> None:
|
|
"""Drain one client's outbound queue."""
|
|
while True:
|
|
packet = await client.queue.get()
|
|
client.writer.write(packet)
|
|
await client.writer.drain()
|
|
try:
|
|
device_id, payload_text = decode_packet_for_log(packet)
|
|
self._log_socket_traffic(
|
|
"Locator socket sent to %s: device_id=%d payload=%s"
|
|
% (client.peer_name, device_id, payload_text)
|
|
)
|
|
except ValueError as error:
|
|
self._log_socket_traffic(
|
|
"Locator socket sent undecodable packet to %s: bytes=%d error=%s"
|
|
% (client.peer_name, len(packet), error)
|
|
)
|
|
|
|
async def _receive_packets(
|
|
self,
|
|
reader: asyncio.StreamReader,
|
|
client: _ClientConnection,
|
|
) -> None:
|
|
"""Receive inbound client packets and queue valid speed updates."""
|
|
while True:
|
|
device_id, payload = await read_packet_with_limit(reader, self._max_payload_bytes)
|
|
payload_text = format_payload_for_log(payload)
|
|
if isinstance(payload, dict) and "vlc" in payload:
|
|
speed_m_s = parse_vlc(payload)
|
|
self._speed_updates.put(speed_m_s)
|
|
self._log_socket_traffic(
|
|
"Locator socket received from %s: device_id=%d payload=%s speed_m_s=%g"
|
|
% (client.peer_name, device_id, payload_text, speed_m_s)
|
|
)
|
|
continue
|
|
|
|
self._log_socket_traffic(
|
|
"Locator socket received from %s: device_id=%d payload=%s"
|
|
% (client.peer_name, device_id, payload_text)
|
|
)
|
|
|
|
def _broadcast_packet(self, packet: bytes) -> None:
|
|
"""Enqueue one packet for all connected clients."""
|
|
for client in list(self._clients):
|
|
self._enqueue_packet(client, packet)
|
|
|
|
def _enqueue_packet(self, client: _ClientConnection, packet: bytes) -> None:
|
|
"""Enqueue one packet or disconnect a backpressured client."""
|
|
if client.closed:
|
|
return
|
|
|
|
try:
|
|
client.queue.put_nowait(packet)
|
|
except asyncio.QueueFull:
|
|
client.closed = True
|
|
self._logger.warning(
|
|
"Disconnecting locator client %s after outbound queue overflow.",
|
|
client.peer_name,
|
|
)
|
|
client.writer.close()
|