394 lines
13 KiB
Python
Executable File
394 lines
13 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
"""Independent ROS 2 supervisor for TG3 teleoperation data collection.
|
|
|
|
This node only controls the project-owned ``ros2 bag record`` process through
|
|
``DataRecorderManager``. It never calls, stops, or reconfigures TianGong's
|
|
factory ``record_bag_node``.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import math
|
|
import sys
|
|
import threading
|
|
import tomllib
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import rclpy
|
|
from rclpy.node import Node
|
|
from rclpy.qos import (
|
|
DurabilityPolicy,
|
|
HistoryPolicy,
|
|
QoSProfile,
|
|
ReliabilityPolicy,
|
|
)
|
|
from rclpy.utilities import remove_ros_args
|
|
from std_msgs.msg import String
|
|
|
|
from data_collection import DataRecorderManager, RecorderConfig
|
|
from data_recorder_protocol import (
|
|
RecorderControlProtocol,
|
|
topics_without_publishers,
|
|
)
|
|
|
|
|
|
DEFAULT_CONFIG = Path(__file__).with_name("config.toml")
|
|
|
|
|
|
def _boolean(value: Any, name: str) -> bool:
|
|
if type(value) is not bool:
|
|
raise ValueError(f"data_collection.{name} must be true or false")
|
|
return value
|
|
|
|
|
|
def _positive_number(value: Any, name: str) -> float:
|
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
raise ValueError(f"data_collection.{name} must be a number")
|
|
result = float(value)
|
|
if not math.isfinite(result) or result <= 0.0:
|
|
raise ValueError(f"data_collection.{name} must be positive and finite")
|
|
return result
|
|
|
|
|
|
def _nonnegative_number(value: Any, name: str) -> float:
|
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
raise ValueError(f"data_collection.{name} must be a number")
|
|
result = float(value)
|
|
if not math.isfinite(result) or result < 0.0:
|
|
raise ValueError(
|
|
f"data_collection.{name} must be non-negative and finite"
|
|
)
|
|
return result
|
|
|
|
|
|
def _topic(value: Any, name: str) -> str:
|
|
if (
|
|
not isinstance(value, str)
|
|
or not value.startswith("/")
|
|
or value.strip() != value
|
|
or any(character.isspace() for character in value)
|
|
):
|
|
raise ValueError(
|
|
f"data_collection.{name} must be an absolute ROS topic name"
|
|
)
|
|
return value
|
|
|
|
|
|
def _topic_list(value: Any, name: str, *, allow_empty: bool) -> tuple[str, ...]:
|
|
if not isinstance(value, list):
|
|
raise ValueError(f"data_collection.{name} must be a TOML array")
|
|
result = tuple(_topic(item, name) for item in value)
|
|
if not allow_empty and not result:
|
|
raise ValueError(f"data_collection.{name} must not be empty")
|
|
if len(set(result)) != len(result):
|
|
raise ValueError(f"data_collection.{name} contains duplicate topics")
|
|
return result
|
|
|
|
|
|
def _minimum_topic_rates(value: Any) -> dict[str, float]:
|
|
name = "minimum_topic_rates_hz"
|
|
if not isinstance(value, dict):
|
|
raise ValueError(f"data_collection.{name} must be a TOML table")
|
|
result: dict[str, float] = {}
|
|
for topic, rate in value.items():
|
|
normalized_topic = _topic(topic, f"{name} key")
|
|
result[normalized_topic] = _positive_number(
|
|
rate, f"{name}.{normalized_topic}"
|
|
)
|
|
return result
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SupervisorSettings:
|
|
enabled: bool
|
|
control_topic: str
|
|
status_topic: str
|
|
heartbeat_timeout_seconds: float
|
|
status_publish_period_seconds: float
|
|
required_topics: tuple[str, ...]
|
|
shutdown_timeout_seconds: float
|
|
recorder: RecorderConfig
|
|
|
|
|
|
def load_settings(path: Path | str) -> SupervisorSettings:
|
|
"""Load only ``[data_collection]`` and tolerate bridge-only extra keys."""
|
|
|
|
config_path = Path(path).expanduser()
|
|
with config_path.open("rb") as stream:
|
|
document = tomllib.load(stream)
|
|
section = document.get("data_collection")
|
|
if not isinstance(section, dict):
|
|
raise ValueError("config.toml is missing [data_collection]")
|
|
|
|
enabled = _boolean(section.get("enabled", True), "enabled")
|
|
control_topic = _topic(
|
|
section.get("control_topic", "/tg3/data_collection/control"),
|
|
"control_topic",
|
|
)
|
|
status_topic = _topic(
|
|
section.get("status_topic", "/tg3/data_collection/status"),
|
|
"status_topic",
|
|
)
|
|
if control_topic == status_topic:
|
|
raise ValueError("data collection control and status topics must differ")
|
|
topics = _topic_list(section.get("topics", []), "topics", allow_empty=False)
|
|
required_topics = _topic_list(
|
|
section.get("required_topics", []),
|
|
"required_topics",
|
|
allow_empty=True,
|
|
)
|
|
unknown_required = sorted(set(required_topics) - set(topics))
|
|
if unknown_required:
|
|
raise ValueError(
|
|
"data_collection.required_topics must also appear in topics: "
|
|
+ ", ".join(unknown_required)
|
|
)
|
|
minimum_topic_rates = _minimum_topic_rates(
|
|
section.get("minimum_topic_rates_hz", {})
|
|
)
|
|
unknown_rate_topics = sorted(
|
|
set(minimum_topic_rates) - set(required_topics)
|
|
)
|
|
if unknown_rate_topics:
|
|
raise ValueError(
|
|
"data_collection.minimum_topic_rates_hz keys must also appear in "
|
|
"required_topics: " + ", ".join(unknown_rate_topics)
|
|
)
|
|
|
|
heartbeat_timeout = _positive_number(
|
|
section.get("heartbeat_timeout_seconds", 3.0),
|
|
"heartbeat_timeout_seconds",
|
|
)
|
|
status_rate = _positive_number(
|
|
section.get("status_publish_rate_hz", 5.0),
|
|
"status_publish_rate_hz",
|
|
)
|
|
minimum_free_gib = _nonnegative_number(
|
|
section.get("minimum_free_gib", 10.0), "minimum_free_gib"
|
|
)
|
|
max_duration = _positive_number(
|
|
section.get("max_duration_seconds", 1800.0),
|
|
"max_duration_seconds",
|
|
)
|
|
base_directory = section.get(
|
|
"base_directory", "/home/nvidia/tg3_data_collection"
|
|
)
|
|
if not isinstance(base_directory, str) or not base_directory.strip():
|
|
raise ValueError(
|
|
"data_collection.base_directory must be a non-empty path"
|
|
)
|
|
|
|
recorder = RecorderConfig(
|
|
base_directory=Path(base_directory).expanduser(),
|
|
topics=topics,
|
|
required_topics=required_topics,
|
|
minimum_topic_rates_hz=minimum_topic_rates,
|
|
minimum_free_bytes=int(minimum_free_gib * 1024**3),
|
|
max_duration_seconds=max_duration,
|
|
poll_interval_seconds=0.1,
|
|
sigint_timeout_seconds=15.0,
|
|
kill_timeout_seconds=3.0,
|
|
ros2_executable="ros2",
|
|
validate_bag_info=True,
|
|
bag_info_timeout_seconds=15.0,
|
|
)
|
|
return SupervisorSettings(
|
|
enabled=enabled,
|
|
control_topic=control_topic,
|
|
status_topic=status_topic,
|
|
heartbeat_timeout_seconds=heartbeat_timeout,
|
|
status_publish_period_seconds=1.0 / status_rate,
|
|
required_topics=required_topics,
|
|
# Finalization includes rosbag2 SIGINT, bag-info validation and SHA256
|
|
# generation. Leave enough time for a large, fully valid episode.
|
|
shutdown_timeout_seconds=_positive_number(
|
|
section.get("shutdown_timeout_seconds", 120.0),
|
|
"shutdown_timeout_seconds",
|
|
),
|
|
recorder=recorder,
|
|
)
|
|
|
|
|
|
class DataRecorderNode(Node):
|
|
"""Reliable String-topic adapter around ``RecorderControlProtocol``."""
|
|
|
|
def __init__(self, settings: SupervisorSettings) -> None:
|
|
super().__init__("tg3_data_recorder_supervisor")
|
|
self.settings = settings
|
|
self._closed = False
|
|
self._close_lock = threading.Lock()
|
|
self._last_logged_error = ""
|
|
|
|
recorder = DataRecorderManager(settings.recorder)
|
|
self.protocol = RecorderControlProtocol(
|
|
recorder,
|
|
heartbeat_timeout_seconds=settings.heartbeat_timeout_seconds,
|
|
start_preflight=self._start_preflight,
|
|
)
|
|
|
|
control_qos = QoSProfile(
|
|
history=HistoryPolicy.KEEP_LAST,
|
|
depth=20,
|
|
reliability=ReliabilityPolicy.RELIABLE,
|
|
durability=DurabilityPolicy.VOLATILE,
|
|
)
|
|
status_qos = QoSProfile(
|
|
history=HistoryPolicy.KEEP_LAST,
|
|
depth=10,
|
|
reliability=ReliabilityPolicy.RELIABLE,
|
|
durability=DurabilityPolicy.TRANSIENT_LOCAL,
|
|
)
|
|
self._status_publisher = self.create_publisher(
|
|
String, settings.status_topic, status_qos
|
|
)
|
|
self._control_subscription = self.create_subscription(
|
|
String,
|
|
settings.control_topic,
|
|
self._on_control,
|
|
control_qos,
|
|
)
|
|
self._timer = self.create_timer(
|
|
settings.status_publish_period_seconds, self._on_timer
|
|
)
|
|
self._publish_status(self.protocol.status())
|
|
self.get_logger().info(
|
|
"TG3 data recorder ready: control=%s status=%s base=%s topics=%d "
|
|
"heartbeat_timeout=%.3fs"
|
|
% (
|
|
settings.control_topic,
|
|
settings.status_topic,
|
|
settings.recorder.base_directory,
|
|
len(settings.recorder.topics),
|
|
settings.heartbeat_timeout_seconds,
|
|
)
|
|
)
|
|
|
|
def _start_preflight(self) -> tuple[bool, str]:
|
|
if not self.settings.required_topics:
|
|
return True, ""
|
|
try:
|
|
missing = topics_without_publishers(
|
|
self.settings.required_topics,
|
|
self.get_publishers_info_by_topic,
|
|
)
|
|
except Exception as error:
|
|
return False, f"cannot inspect ROS graph: {type(error).__name__}: {error}"
|
|
if missing:
|
|
return (
|
|
False,
|
|
"required ROS topics have no live publisher: "
|
|
+ ", ".join(missing),
|
|
)
|
|
return True, ""
|
|
|
|
def _on_control(self, message: String) -> None:
|
|
try:
|
|
status = self.protocol.handle_json(message.data)
|
|
self._publish_status(status)
|
|
except Exception as error:
|
|
# A malformed or otherwise bad data request must never kill the
|
|
# node or have any effect on the independent teleoperation bridge.
|
|
self.get_logger().error(
|
|
f"data recorder control callback failed: {type(error).__name__}: {error}"
|
|
)
|
|
|
|
def _on_timer(self) -> None:
|
|
try:
|
|
self._publish_status(self.protocol.poll())
|
|
except Exception as error:
|
|
self.get_logger().error(
|
|
f"data recorder status timer failed: {type(error).__name__}: {error}"
|
|
)
|
|
|
|
def _publish_status(self, status: dict[str, Any]) -> None:
|
|
message = String()
|
|
message.data = json.dumps(
|
|
status,
|
|
ensure_ascii=False,
|
|
sort_keys=True,
|
|
separators=(",", ":"),
|
|
allow_nan=False,
|
|
)
|
|
self._status_publisher.publish(message)
|
|
error = status.get("last_error")
|
|
error_text = error if isinstance(error, str) else ""
|
|
if error_text and error_text != self._last_logged_error:
|
|
self.get_logger().warning(f"data recorder status error: {error_text}")
|
|
self._last_logged_error = error_text
|
|
|
|
def close(self) -> bool:
|
|
with self._close_lock:
|
|
if self._closed:
|
|
return True
|
|
self._closed = True
|
|
self.get_logger().info("stopping data recorder supervisor")
|
|
try:
|
|
return self.protocol.shutdown(
|
|
timeout=self.settings.shutdown_timeout_seconds
|
|
)
|
|
except Exception as error:
|
|
self.get_logger().error(
|
|
f"data recorder shutdown failed: {type(error).__name__}: {error}"
|
|
)
|
|
return False
|
|
|
|
|
|
def _arguments(argv: list[str]) -> argparse.Namespace:
|
|
parser = argparse.ArgumentParser(
|
|
description="TG3 independent ROS 2 MCAP recorder supervisor"
|
|
)
|
|
parser.add_argument(
|
|
"--config",
|
|
type=Path,
|
|
default=DEFAULT_CONFIG,
|
|
help="TG3 bridge config.toml containing [data_collection]",
|
|
)
|
|
return parser.parse_args(remove_ros_args(argv)[1:])
|
|
|
|
|
|
def main(argv: list[str] | None = None) -> int:
|
|
process_argv = sys.argv if argv is None else argv
|
|
arguments = _arguments(process_argv)
|
|
try:
|
|
settings = load_settings(arguments.config)
|
|
except Exception as error:
|
|
print(
|
|
f"data recorder configuration error: {type(error).__name__}: {error}",
|
|
file=sys.stderr,
|
|
)
|
|
return 2
|
|
if not settings.enabled:
|
|
print("TG3 data collection is disabled in config.toml")
|
|
return 0
|
|
|
|
rclpy.init(args=process_argv)
|
|
node: DataRecorderNode | None = None
|
|
exit_code = 0
|
|
try:
|
|
node = DataRecorderNode(settings)
|
|
rclpy.spin(node)
|
|
except KeyboardInterrupt:
|
|
pass
|
|
except Exception as error:
|
|
print(
|
|
f"data recorder supervisor failed: {type(error).__name__}: {error}",
|
|
file=sys.stderr,
|
|
)
|
|
exit_code = 1
|
|
finally:
|
|
if node is not None:
|
|
if not node.close():
|
|
exit_code = 1
|
|
node.destroy_node()
|
|
if rclpy.ok():
|
|
rclpy.shutdown()
|
|
return exit_code
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|