GrabBag/Test/BeltTearingTcpProtocolTest/belt_tearing_tcp_test.py

483 lines
20 KiB
Python

#!/usr/bin/env python3
"""GUI client for validating the BeltTearing legacy TCP point-cloud stream."""
from __future__ import annotations
import argparse
import csv
import json
import queue
import socket
import struct
import threading
import time
from dataclasses import dataclass
from pathlib import Path
import tkinter as tk
from tkinter import filedialog, messagebox, ttk
TYPE_TEXT = 0x01
TYPE_IMAGE = 0x02
TYPE_READ_CONFIG = 0x03
TYPE_WRITE_CONFIG = 0x04
TYPE_POINT_CLOUD = 0x05
FRAME_TAIL = b"___END___\r\n"
MAX_PAYLOAD_SIZE = 256 * 1024 * 1024
POINT_CLOUD_HEADER = struct.Struct(">BQQIIII")
POINT_RECORD = struct.Struct(">iddd")
@dataclass(frozen=True)
class PointCloudSummary:
frame_index: int
timestamp: int
original_count: int
sampled_count: int
line_step: int
point_step: int
x_range: tuple[float, float]
y_range: tuple[float, float]
z_range: tuple[float, float]
def receive_exact(sock: socket.socket, size: int) -> bytes:
chunks: list[bytes] = []
remaining = size
while remaining:
data = sock.recv(remaining)
if not data:
raise ConnectionError("The server closed the connection")
chunks.append(data)
remaining -= len(data)
return b"".join(chunks)
def receive_frame(sock: socket.socket) -> tuple[int, bytes, int]:
header = receive_exact(sock, 5)
data_type, payload_size = struct.unpack(">BI", header)
if payload_size > MAX_PAYLOAD_SIZE:
raise ValueError(f"Payload is too large: {payload_size} bytes")
payload = receive_exact(sock, payload_size)
tail = receive_exact(sock, len(FRAME_TAIL))
if tail != FRAME_TAIL:
raise ValueError(f"Invalid frame tail: {tail!r}")
return data_type, payload, len(header) + len(payload) + len(tail)
def build_frame(data_type: int, payload: bytes = b"") -> bytes:
return struct.pack(">BI", data_type, len(payload)) + payload + FRAME_TAIL
def parse_point_cloud(payload: bytes) -> tuple[PointCloudSummary, list[tuple[int, float, float, float]]]:
if len(payload) < POINT_CLOUD_HEADER.size:
raise ValueError(f"Point-cloud payload is too short: {len(payload)} bytes")
(format_version, frame_index, timestamp, original_count, sampled_count,
line_step, point_step) = POINT_CLOUD_HEADER.unpack_from(payload)
if format_version != 1:
raise ValueError(f"Unsupported point-cloud payload version: {format_version}")
expected_size = POINT_CLOUD_HEADER.size + sampled_count * POINT_RECORD.size
if len(payload) != expected_size:
raise ValueError(
f"Point-cloud size mismatch: expected {expected_size}, received {len(payload)}"
)
points: list[tuple[int, float, float, float]] = []
offset = POINT_CLOUD_HEADER.size
for _ in range(sampled_count):
points.append(POINT_RECORD.unpack_from(payload, offset))
offset += POINT_RECORD.size
if points:
xs = [point[1] for point in points]
ys = [point[2] for point in points]
zs = [point[3] for point in points]
x_range = (min(xs), max(xs))
y_range = (min(ys), max(ys))
z_range = (min(zs), max(zs))
else:
x_range = y_range = z_range = (0.0, 0.0)
summary = PointCloudSummary(
frame_index=frame_index,
timestamp=timestamp,
original_count=original_count,
sampled_count=sampled_count,
line_step=line_step,
point_step=point_step,
x_range=x_range,
y_range=y_range,
z_range=z_range,
)
return summary, points
class BeltTearingTcpTestApp:
def __init__(self, root: tk.Tk, host: str, port: int) -> None:
self.root = root
self.root.title("BeltTearing TCP Protocol Test")
self.root.geometry("920x680")
self.root.minsize(820, 600)
self.host_var = tk.StringVar(value=host)
self.port_var = tk.StringVar(value=str(port))
self.output_dir_var = tk.StringVar(value=str(Path.cwd() / "tcp_test_output"))
self.save_points_var = tk.BooleanVar(value=False)
self.save_images_var = tk.BooleanVar(value=False)
self.send_enabled_var = tk.BooleanVar(value=True)
self.line_step_var = tk.StringVar(value="1")
self.point_step_var = tk.StringVar(value="1")
self.connection_var = tk.StringVar(value="Disconnected")
self.lines_var = tk.StringVar(value="0")
self.bytes_var = tk.StringVar(value="0")
self.errors_var = tk.StringVar(value="0")
self.frame_var = tk.StringVar(value="-")
self.count_var = tk.StringVar(value="-")
self.steps_var = tk.StringVar(value="-")
self.x_range_var = tk.StringVar(value="-")
self.y_range_var = tk.StringVar(value="-")
self.z_range_var = tk.StringVar(value="-")
self.sock: socket.socket | None = None
self.socket_lock = threading.Lock()
self.stop_event = threading.Event()
self.receiver_thread: threading.Thread | None = None
self.events: queue.Queue[tuple[str, object]] = queue.Queue()
self.point_line_count = 0
self.received_bytes = 0
self.error_count = 0
self.image_count = 0
self._build_ui()
self.root.protocol("WM_DELETE_WINDOW", self.close)
self.root.after(100, self._drain_events)
def _build_ui(self) -> None:
connection = ttk.LabelFrame(self.root, text="Connection", padding=10)
connection.pack(fill=tk.X, padx=10, pady=(10, 5))
ttk.Label(connection, text="Host").grid(row=0, column=0, sticky=tk.W)
ttk.Entry(connection, textvariable=self.host_var, width=20).grid(row=0, column=1, padx=6)
ttk.Label(connection, text="Port").grid(row=0, column=2, sticky=tk.W)
ttk.Entry(connection, textvariable=self.port_var, width=8).grid(row=0, column=3, padx=6)
self.connect_button = ttk.Button(connection, text="Connect", command=self.connect)
self.connect_button.grid(row=0, column=4, padx=6)
self.disconnect_button = ttk.Button(
connection, text="Disconnect", command=self.disconnect, state=tk.DISABLED
)
self.disconnect_button.grid(row=0, column=5, padx=6)
self.read_config_button = ttk.Button(
connection, text="Read Config", command=self.read_config, state=tk.DISABLED
)
self.read_config_button.grid(row=0, column=6, padx=6)
ttk.Label(connection, textvariable=self.connection_var).grid(
row=0, column=7, padx=(20, 0), sticky=tk.W
)
ttk.Checkbutton(
connection, text="Point cloud enabled", variable=self.send_enabled_var
).grid(row=1, column=0, columnspan=2, pady=(10, 0), sticky=tk.W)
ttk.Label(connection, text="Line step").grid(row=1, column=2, pady=(10, 0), sticky=tk.E)
ttk.Entry(connection, textvariable=self.line_step_var, width=8).grid(
row=1, column=3, padx=6, pady=(10, 0)
)
ttk.Label(connection, text="Point step").grid(row=1, column=4, pady=(10, 0), sticky=tk.E)
ttk.Entry(connection, textvariable=self.point_step_var, width=8).grid(
row=1, column=5, padx=6, pady=(10, 0)
)
self.set_sampling_button = ttk.Button(
connection, text="Set Sampling", command=self.set_sampling, state=tk.DISABLED
)
self.set_sampling_button.grid(row=1, column=6, padx=6, pady=(10, 0))
storage = ttk.LabelFrame(self.root, text="Optional Capture", padding=10)
storage.pack(fill=tk.X, padx=10, pady=5)
ttk.Checkbutton(storage, text="Append point cloud to CSV", variable=self.save_points_var).grid(
row=0, column=0, sticky=tk.W
)
ttk.Checkbutton(storage, text="Save JPEG frames", variable=self.save_images_var).grid(
row=0, column=1, padx=16, sticky=tk.W
)
ttk.Entry(storage, textvariable=self.output_dir_var).grid(
row=1, column=0, columnspan=2, pady=(8, 0), sticky=tk.EW
)
ttk.Button(storage, text="Browse", command=self.choose_output_dir).grid(
row=1, column=2, padx=(8, 0), pady=(8, 0)
)
storage.columnconfigure(1, weight=1)
stats = ttk.LabelFrame(self.root, text="Point Cloud", padding=10)
stats.pack(fill=tk.X, padx=10, pady=5)
labels = (
("Received lines", self.lines_var),
("Received bytes", self.bytes_var),
("Protocol errors", self.errors_var),
("Last frame", self.frame_var),
("Points original / sampled", self.count_var),
("Line step / point step", self.steps_var),
("X range", self.x_range_var),
("Y range", self.y_range_var),
("Z range", self.z_range_var),
)
for index, (label, value) in enumerate(labels):
row, column = divmod(index, 3)
cell = ttk.Frame(stats)
cell.grid(row=row, column=column, padx=8, pady=5, sticky=tk.EW)
ttk.Label(cell, text=label).pack(anchor=tk.W)
ttk.Label(cell, textvariable=value, font=("Consolas", 11, "bold")).pack(anchor=tk.W)
for column in range(3):
stats.columnconfigure(column, weight=1)
log_frame = ttk.LabelFrame(self.root, text="Protocol Log", padding=8)
log_frame.pack(fill=tk.BOTH, expand=True, padx=10, pady=(5, 10))
self.log_text = tk.Text(log_frame, wrap=tk.NONE, state=tk.DISABLED, font=("Consolas", 10))
y_scroll = ttk.Scrollbar(log_frame, orient=tk.VERTICAL, command=self.log_text.yview)
x_scroll = ttk.Scrollbar(log_frame, orient=tk.HORIZONTAL, command=self.log_text.xview)
self.log_text.configure(yscrollcommand=y_scroll.set, xscrollcommand=x_scroll.set)
self.log_text.grid(row=0, column=0, sticky=tk.NSEW)
y_scroll.grid(row=0, column=1, sticky=tk.NS)
x_scroll.grid(row=1, column=0, sticky=tk.EW)
log_frame.rowconfigure(0, weight=1)
log_frame.columnconfigure(0, weight=1)
def connect(self) -> None:
try:
port = int(self.port_var.get())
if not 1 <= port <= 65535:
raise ValueError
except ValueError:
messagebox.showerror("Invalid port", "Port must be in the range 1-65535.")
return
self.disconnect()
self.stop_event.clear()
self.connection_var.set("Connecting...")
self.connect_button.configure(state=tk.DISABLED)
self.receiver_thread = threading.Thread(
target=self._receiver_main,
args=(self.host_var.get().strip(), port),
daemon=True,
)
self.receiver_thread.start()
def disconnect(self) -> None:
self.stop_event.set()
with self.socket_lock:
sock, self.sock = self.sock, None
if sock is not None:
try:
sock.shutdown(socket.SHUT_RDWR)
except OSError:
pass
sock.close()
receiver_thread = self.receiver_thread
if (receiver_thread is not None and receiver_thread.is_alive() and
receiver_thread is not threading.current_thread()):
receiver_thread.join(timeout=1.0)
self.receiver_thread = None
self._set_connected(False)
def read_config(self) -> None:
try:
self._send(build_frame(TYPE_READ_CONFIG))
self._log("ReadConfig request sent")
except OSError as error:
self._log(f"ReadConfig failed: {error}")
def set_sampling(self) -> None:
try:
line_step = int(self.line_step_var.get())
point_step = int(self.point_step_var.get())
if line_step <= 0 or point_step <= 0:
raise ValueError
except ValueError:
messagebox.showerror("Invalid sampling", "Line step and point step must be positive integers.")
return
request = {
"command": "setAlgorithmParams",
"pointCloudEnabled": self.send_enabled_var.get(),
"pointCloudLineStep": line_step,
"pointCloudPointStep": point_step,
}
try:
payload = json.dumps(request, separators=(",", ":")).encode("utf-8")
self._send(build_frame(TYPE_WRITE_CONFIG, payload))
self._log(
f"Point-cloud sampling sent: enabled={self.send_enabled_var.get()} "
f"lineStep={line_step} pointStep={point_step}"
)
self.root.after(300, self.read_config)
except OSError as error:
self._log(f"Set sampling failed: {error}")
def choose_output_dir(self) -> None:
selected = filedialog.askdirectory(initialdir=self.output_dir_var.get())
if selected:
self.output_dir_var.set(selected)
def close(self) -> None:
self.disconnect()
self.root.destroy()
def _receiver_main(self, host: str, port: int) -> None:
try:
sock = socket.create_connection((host, port), timeout=5.0)
sock.settimeout(None)
with self.socket_lock:
if self.stop_event.is_set():
sock.close()
return
self.sock = sock
self.events.put(("connected", f"Connected to {host}:{port}"))
while not self.stop_event.is_set():
data_type, payload, frame_size = receive_frame(sock)
self.events.put(("frame", (data_type, payload, frame_size)))
except (ConnectionError, OSError, ValueError) as error:
if not self.stop_event.is_set():
self.events.put(("error", str(error)))
finally:
with self.socket_lock:
if self.sock is not None:
self.sock.close()
self.sock = None
self.events.put(("disconnected", "Disconnected"))
def _send(self, data: bytes) -> None:
with self.socket_lock:
if self.sock is None:
raise OSError("Not connected")
self.sock.sendall(data)
def _drain_events(self) -> None:
try:
while True:
event, value = self.events.get_nowait()
if event == "connected":
self._set_connected(True)
self._log(str(value))
elif event == "disconnected":
self._set_connected(False)
self._log(str(value))
elif event == "error":
self.error_count += 1
self.errors_var.set(str(self.error_count))
self._log(f"ERROR: {value}")
elif event == "frame":
data_type, payload, frame_size = value # type: ignore[misc]
self._handle_frame(data_type, payload, frame_size)
except queue.Empty:
pass
if self.root.winfo_exists():
self.root.after(100, self._drain_events)
def _handle_frame(self, data_type: int, payload: bytes, frame_size: int) -> None:
self.received_bytes += frame_size
self.bytes_var.set(f"{self.received_bytes:,}")
try:
if data_type == TYPE_POINT_CLOUD:
summary, points = parse_point_cloud(payload)
self.point_line_count += 1
self.lines_var.set(str(self.point_line_count))
self.frame_var.set(f"{summary.frame_index} (ts={summary.timestamp})")
self.count_var.set(f"{summary.original_count} / {summary.sampled_count}")
self.steps_var.set(f"{summary.line_step} / {summary.point_step}")
self.x_range_var.set(self._format_range(summary.x_range))
self.y_range_var.set(self._format_range(summary.y_range))
self.z_range_var.set(self._format_range(summary.z_range))
if self.point_line_count <= 5 or self.point_line_count % 100 == 0:
self._log(
f"PointCloud frame={summary.frame_index} points="
f"{summary.original_count}->{summary.sampled_count} "
f"steps={summary.line_step}/{summary.point_step}"
)
if self.save_points_var.get():
self._append_points(summary, points)
elif data_type == TYPE_IMAGE:
self.image_count += 1
if self.save_images_var.get():
output_dir = self._ensure_output_dir()
path = output_dir / f"image_{self.image_count:06d}.jpg"
path.write_bytes(payload)
if self.image_count <= 3 or self.image_count % 20 == 0:
self._log(f"Image #{self.image_count}: {len(payload):,} bytes")
elif data_type in (TYPE_TEXT, TYPE_READ_CONFIG, TYPE_WRITE_CONFIG):
self._handle_json(data_type, payload)
else:
raise ValueError(f"Unknown data type: 0x{data_type:02X}")
except (OSError, UnicodeError, ValueError, json.JSONDecodeError) as error:
self.error_count += 1
self.errors_var.set(str(self.error_count))
self._log(f"ERROR type=0x{data_type:02X}: {error}")
def _handle_json(self, data_type: int, payload: bytes) -> None:
decoded = json.loads(payload.decode("utf-8"))
if isinstance(decoded, dict):
config = decoded.get("config")
if isinstance(config, dict):
point_cloud = config.get("pointCloudSendParam")
if isinstance(point_cloud, dict):
self.send_enabled_var.set(bool(point_cloud.get("enabled", True)))
self.line_step_var.set(str(point_cloud.get("lineStep", 1)))
self.point_step_var.set(str(point_cloud.get("pointStep", 1)))
type_name = {
TYPE_TEXT: "Text",
TYPE_READ_CONFIG: "ReadConfig",
TYPE_WRITE_CONFIG: "WriteConfig",
}[data_type]
self._log(f"{type_name}: {json.dumps(decoded, ensure_ascii=False, separators=(',', ':'))}")
def _append_points(
self,
summary: PointCloudSummary,
points: list[tuple[int, float, float, float]],
) -> None:
path = self._ensure_output_dir() / "point_cloud.csv"
write_header = not path.exists() or path.stat().st_size == 0
with path.open("a", newline="", encoding="utf-8") as csv_file:
writer = csv.writer(csv_file)
if write_header:
writer.writerow(("frame_index", "timestamp", "point_index", "x", "y", "z"))
for point_index, x, y, z in points:
writer.writerow((summary.frame_index, summary.timestamp, point_index, x, y, z))
def _ensure_output_dir(self) -> Path:
output_dir = Path(self.output_dir_var.get()).expanduser()
output_dir.mkdir(parents=True, exist_ok=True)
return output_dir
def _set_connected(self, connected: bool) -> None:
self.connection_var.set("Connected" if connected else "Disconnected")
self.connect_button.configure(state=tk.DISABLED if connected else tk.NORMAL)
self.disconnect_button.configure(state=tk.NORMAL if connected else tk.DISABLED)
self.read_config_button.configure(state=tk.NORMAL if connected else tk.DISABLED)
self.set_sampling_button.configure(state=tk.NORMAL if connected else tk.DISABLED)
def _log(self, message: str) -> None:
timestamp = time.strftime("%H:%M:%S")
self.log_text.configure(state=tk.NORMAL)
self.log_text.insert(tk.END, f"[{timestamp}] {message}\n")
self.log_text.see(tk.END)
self.log_text.configure(state=tk.DISABLED)
@staticmethod
def _format_range(value_range: tuple[float, float]) -> str:
return f"{value_range[0]:.3f} .. {value_range[1]:.3f}"
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", default="127.0.0.1", help="BeltTearing server address")
parser.add_argument("--port", type=int, default=5900, help="Legacy binary protocol port")
args = parser.parse_args()
root = tk.Tk()
BeltTearingTcpTestApp(root, args.host, args.port)
root.mainloop()
if __name__ == "__main__":
main()