From 6caad268b866cc623aa411791952f06265950fbb Mon Sep 17 00:00:00 2001 From: LengedZhao <867178694@qq.com> Date: Mon, 10 Aug 2026 15:54:15 +0800 Subject: [PATCH] feat: add session-gated TG3 data collection --- Data_Get/.gitignore | 3 + Data_Get/README.md | 8 + README.md | 21 +- docs/天工3.0本地同构臂遥操迁移部署指南.md | 86 ++ tg3_data_collection/README.md | 73 ++ tg3_data_collection/data_get_sync.py | 363 ++++++ tg3_data_collection/test_data_get_sync.py | 149 +++ tg3_data_collection/tg3-data-get-sync.service | 14 + tg3_local_teleop/README.md | 59 +- tg3_local_teleop/config.toml | 67 +- tg3_local_teleop/data_collection.py | 1140 +++++++++++++++++ tg3_local_teleop/data_recorder_node.py | 368 ++++++ tg3_local_teleop/data_recorder_protocol.py | 546 ++++++++ tg3_local_teleop/run_data_recorder.sh | 26 + tg3_local_teleop/test_data_collection.py | 699 ++++++++++ .../test_data_recorder_protocol.py | 274 ++++ tg3_local_teleop/test_session_gate.py | 154 +++ tg3_local_teleop/tg3-data-recorder.service | 23 + tg3_local_teleop/tg3-local-teleop.service | 4 +- tg3_local_teleop/tg3_local_teleop.py | 345 ++++- .../tg3-omnisocket-sender.service | 2 +- verify.sh | 14 +- 22 files changed, 4429 insertions(+), 9 deletions(-) create mode 100644 Data_Get/.gitignore create mode 100644 Data_Get/README.md create mode 100644 tg3_data_collection/README.md create mode 100755 tg3_data_collection/data_get_sync.py create mode 100755 tg3_data_collection/test_data_get_sync.py create mode 100644 tg3_data_collection/tg3-data-get-sync.service create mode 100755 tg3_local_teleop/data_collection.py create mode 100755 tg3_local_teleop/data_recorder_node.py create mode 100755 tg3_local_teleop/data_recorder_protocol.py create mode 100755 tg3_local_teleop/run_data_recorder.sh create mode 100755 tg3_local_teleop/test_data_collection.py create mode 100755 tg3_local_teleop/test_data_recorder_protocol.py create mode 100644 tg3_local_teleop/tg3-data-recorder.service diff --git a/Data_Get/.gitignore b/Data_Get/.gitignore new file mode 100644 index 0000000..7c9d611 --- /dev/null +++ b/Data_Get/.gitignore @@ -0,0 +1,3 @@ +* +!.gitignore +!README.md diff --git a/Data_Get/README.md b/Data_Get/README.md new file mode 100644 index 0000000..7062b9d --- /dev/null +++ b/Data_Get/README.md @@ -0,0 +1,8 @@ +# Data_Get + +完成的数采 episode 自动保存到本目录,每个子目录包含 MCAP、`metadata.yaml`、录制日志和 +`manifest.json`、`READY`。数据来自机器人 Nvidia 上独立的项目录制服务,不影响厂家 +`record_bag_node`。运行数据、临时同步目录及 `sync_status.json` 均不会提交到 Git。 + +只有通过大小和 SHA-256 校验的目录才会从 `.incoming` 原子移动到这里。机器人端保留 +ready 副本,传输中断后会自动续传,不要把 `.incoming` 当作完整数据集。 diff --git a/README.md b/README.md index 2f9ac88..db997bd 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,10 @@ 读取数据,仅在左 Z + 右 C 连续 3 秒开启后注册并发送,STOP 后立即断开;还包含 不走公网时在 EAI 运行 `kcpserver` 的本地 Hub 用户服务模板。 - `tg3_local_teleop/`:部署到机器人 Nvidia;直接接收 OmniSocket 数据,发布双臂、 - 双手及 `/hric/robot/cmd_vel`,并提供限速回 Home。 + 双手及 `/hric/robot/cmd_vel`,提供限速回 Home,并独立录制控制数据 MCAP。 +- `tg3_data_collection/`:部署在本机;通过本地 SSH/rsync 拉取已完整收尾并校验的 + episode,原子保存到 `Data_Get/`。不会删除机器人上的备份。 +- `Data_Get/`:最终数采目录。实际 MCAP、清单和同步状态默认不提交 Git。 - `docs/`:可提交的跨机器人迁移步骤。含实测帧和现场拓扑的汇报/证据文档只保留在 当前本地工作副本,不同步到匿名可读的远端仓库。 @@ -44,6 +47,22 @@ TOML 解析检查。若还要核对外置 OmniSocketGo 版本,可设置任务 - 左 Z + 右摇杆左右:HBWALK 原地转向。 - 右 B 连续 1 秒:右手进入厂商“单食指”姿态;松开至少 0.5 秒后再次连续 1 秒退出。手指动作继续受 `400 units/s` 限速。 +- 左摇杆按下(L3)连续 1 秒:开始数采;松开至少 0.5 秒后再次连续 1 秒: + 结束并保存。只有遥操已开启时才接受;Z+C 结束遥操或安全解除会自动结束数采。 + +## 数采保存 + +Nvidia 独立服务录制明确白名单内的机器人实测状态、双臂/双手命令与反馈、行走、 +IMU、电源状态和 xTELE 原始应用帧。它不接管或停止厂家 `/record_bag_node`,默认也不录 +相机和点云。正常结束后先校验 MCAP、生成 SHA-256 `manifest.json` 和 `READY`,再由 +本机服务拉取到: + +```text +/home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/Data_Get// +``` + +传输中断只会留在 `Data_Get/.incoming/`,不会显示为完成 episode;机器人端 ready +副本不会自动删除。部署、状态检查和恢复步骤见 `tg3_data_collection/README.md`。 ## 迁移前必须修改 diff --git a/docs/天工3.0本地同构臂遥操迁移部署指南.md b/docs/天工3.0本地同构臂遥操迁移部署指南.md index 561d3ec..2546470 100644 --- a/docs/天工3.0本地同构臂遥操迁移部署指南.md +++ b/docs/天工3.0本地同构臂遥操迁移部署指南.md @@ -1012,3 +1012,89 @@ Hub IP/端口变化时,额外同时修改 EAI `--server` 和机器人 `omnisoc 当前没有直连断流门控的版本中设置 `locomotion.enabled=false`,同时停用 EAI 的 `tg3-omnisocket-sender.service`,并仅以机器人 monitor-only 模式验证;完成第 7.5 节的 ZMQ 专用断流门控前不得启用运动发布。 + +## 11. 数采部署与迁移 + +数采是本项目新增功能,不调用 `/bag_record/control/notify`,也不停止或修改 Ubuntu 厂家 +`record_bag_node`。Nvidia 独立录制,PS 本机负责把完成数据同步到: + +```text +/home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/Data_Get +``` + +### 11.1 机器人 Nvidia + +部署 `tg3_local_teleop/` 全目录,至少必须包含: + +```text +tg3_local_teleop.py +data_collection.py +data_recorder_protocol.py +data_recorder_node.py +config.toml +run.sh +run_data_recorder.sh +wait_ros_ready.sh +tg3-local-teleop.service +tg3-data-recorder.service +ros2_py/ +``` + +目标路径固定为 `/home/nvidia/tg3_local_teleop`。先按原项目步骤构建 `ros2_py`,再安装: + +```bash +mkdir -p ~/.config/systemd/user +cp /home/nvidia/tg3_local_teleop/tg3-data-recorder.service \ + /home/nvidia/tg3_local_teleop/tg3-local-teleop.service \ + ~/.config/systemd/user/ +systemctl --user daemon-reload +systemctl --user enable --now tg3-data-recorder.service +systemctl --user restart tg3-local-teleop.service +``` + +录制脚本必须能依次加载 `/opt/ros/jazzy`、`/home/nvidia/xos`、 +`/opt/robot_tele_server/install` 和本项目 `ros2_py/install`;否则自定义消息可能无法解码。 +还需 Python 3 的 PyYAML(Ubuntu 包 `python3-yaml`)解析 rosbag metadata。 +迁移时在 `config.toml [data_collection]` 核对 `base_directory`、20 GiB 余量、30 分钟上限、 +明确 topic 白名单和 required topics。不要改成 `ros2 bag record -a`,相机/点云需另行估算 +带宽和磁盘后再加入。 + +### 11.2 PS 本机 + +项目需包含 `tg3_data_collection/` 和 `Data_Get/`。确认到新 Nvidia 的免密 SSH 后安装: + +```bash +ssh -o BatchMode=yes nvidia@192.168.41.2 true +mkdir -p ~/.config/systemd/user +cp tg3_data_collection/tg3-data-get-sync.service ~/.config/systemd/user/ +systemctl --user daemon-reload +systemctl --user enable --now tg3-data-get-sync.service +``` + +机器人 SSH IP 改变时,只改同步 unit;最简命令(替换两个地址)是: + +```bash +sed -i 's/nvidia@旧IP/nvidia@新IP/g' \ + ~/.config/systemd/user/tg3-data-get-sync.service +systemctl --user daemon-reload +systemctl --user restart tg3-data-get-sync.service +``` + +这不修改 OmniSocket Hub、Peer ID 或机器人 `config.toml`。若 PS 用户名/项目路径改变, +同时修改 unit 的 `WorkingDirectory`、脚本路径和 `--destination`。 + +### 11.3 操作与验收 + +1. 保持 L3 松开,长按 Z+C 3 秒正常开启遥操; +2. 左摇杆按下(L3)连续 1 秒开始数采;该键是 `button_joystick.left`,不是 X/Y/Z; +3. 松开 L3 至少 0.5 秒,再长按 1 秒结束;或正常 Z+C 结束遥操自动收尾; +4. 查看 Nvidia `/tg3/data_collection/status`,必须先到 `ready`,不能从 `active/` 取数据; + recorder 会同时要求 required topics 有实时发布者且最终消息计数大于零; +5. 查看本机 `Data_Get/sync_status.json`,最终目录必须有 `READY`、`manifest.json`、 + `bag/metadata.yaml` 和非空 `*.mcap`; +6. 执行 `ros2 bag info /bag` 并确认所需 topic 有消息; +7. 断开 PS 网络再采一条,确认 Nvidia 保留数据;恢复网络后应续传且 SHA-256 通过。 + +每个新会话先要求 L3 稳定松开 0.5 秒,同一次持续按压只切换一次。数采故障不得解除 +遥操或延迟 STOP/Home;桥心跳中断超过 3 秒时 recorder 会自行 SIGINT 收尾。机器人 +`ready/` 不自动删除,确认本机和外部备份后才按具体 episode 清理。 diff --git a/tg3_data_collection/README.md b/tg3_data_collection/README.md new file mode 100644 index 0000000..59ef92b --- /dev/null +++ b/tg3_data_collection/README.md @@ -0,0 +1,73 @@ +# TG3 data episode sync + +机器人只在 `/home/nvidia/tg3_data_collection/ready` 暴露已经收到 SIGINT、写完 +`metadata.yaml`、通过 `ros2 bag info` 并生成 SHA-256 清单的 episode。本机服务用免密 SSH/rsync 复制到 +`Data_Get/.incoming`,逐文件校验后再原子改名为 `Data_Get/`。中断的复制不会 +显示成完成数据,也不会删除机器人上的备份。 + +数据只走本机到 Nvidia 的 SSH 链路,与公网或 EAI 本地 OmniSocket Hub 无关。当前默认: + +```text +nvidia@192.168.41.2:/home/nvidia/tg3_data_collection/ready/ + -> /home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/Data_Get/ +``` + +先确认免密与依赖: + +```bash +ssh -o BatchMode=yes nvidia@192.168.41.2 true +command -v rsync +ssh nvidia@192.168.41.2 command -v rsync +``` + +安装本机用户服务: + +```bash +mkdir -p ~/.config/systemd/user +cp tg3_data_collection/tg3-data-get-sync.service ~/.config/systemd/user/ +systemctl --user daemon-reload +systemctl --user enable --now tg3-data-get-sync.service +``` + +状态: + +```bash +systemctl --user --no-pager status tg3-data-get-sync.service +python3 -m json.tool Data_Get/sync_status.json +``` + +也可以只执行一次,便于首次部署验收: + +```bash +python3 tg3_data_collection/data_get_sync.py --once +``` + +日常轮询不会每 2 秒重新读取并哈希全部历史 MCAP;首次原子发布时已经完成深度校验。 +需要定期审计已有本地数据时单独执行(可能耗时较长): + +```bash +python3 tg3_data_collection/data_get_sync.py --once --verify-existing +``` + +某个旧 episode 损坏会写入 `sync_status.json`,但不会阻止同一轮继续拉取其他新 episode。 + +一个完成目录至少包含: + +```text +/ + READY + manifest.json + bag/metadata.yaml + bag/*.mcap + ros2_bag.stdout.log + ros2_bag.stderr.log +``` + +同步端会再次核对 `READY`、episode ID、每个 MCAP/metadata 的长度和 SHA-256;远端 +manifest 在传输中发生变化也会拒绝发布最终目录。`.incoming`、`sync_status.json` 和全部 +episode 已由 `Data_Get/.gitignore` 排除。服务绝不自动删除机器人端 `ready/` 数据;确认 +本机备份后如需清理,必须由操作者明确指定具体 episode,不能删除整个项目目录。 + +迁移后需要同步修改 service 中的本机项目路径、机器人 SSH 地址和机器人 ready 路径。 +如果 Nvidia 的 `192.168.41.2` 改变,只改该 unit 的 `--remote`;它不在机器人 +`config.toml` 或 OmniSocket Peer 设置中。 diff --git a/tg3_data_collection/data_get_sync.py b/tg3_data_collection/data_get_sync.py new file mode 100755 index 0000000..b889a33 --- /dev/null +++ b/tg3_data_collection/data_get_sync.py @@ -0,0 +1,363 @@ +#!/usr/bin/env python3 +"""Pull completed TG3 data-collection episodes from the robot atomically.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import re +import shlex +import shutil +import signal +import subprocess +import time +from pathlib import Path +from typing import Any, Sequence + + +SAFE_EPISODE_NAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,127}$") + + +def atomic_write_json(path: Path, payload: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text( + json.dumps(payload, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for block in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def safe_episode_name(value: str) -> bool: + return bool(SAFE_EPISODE_NAME.fullmatch(value)) and value not in (".", "..") + + +def safe_relative_path(value: object) -> Path: + if not isinstance(value, str) or not value: + raise ValueError("manifest file path must be a non-empty string") + path = Path(value) + if path.is_absolute() or ".." in path.parts: + raise ValueError(f"unsafe manifest file path: {value!r}") + return path + + +def parse_manifest(raw: str, expected_episode: str) -> dict[str, Any]: + payload = json.loads(raw) + if not isinstance(payload, dict): + raise ValueError("manifest root must be an object") + if payload.get("state") != "complete": + raise ValueError("remote episode is not complete") + if payload.get("episode_id") != expected_episode: + raise ValueError("manifest episode_id does not match directory name") + files = payload.get("files") + if not isinstance(files, list) or not files: + raise ValueError("manifest files must be a non-empty list") + saw_mcap = False + for entry in files: + if not isinstance(entry, dict): + raise ValueError("manifest file entry must be an object") + relative = safe_relative_path(entry.get("path")) + saw_mcap = saw_mcap or relative.suffix == ".mcap" + size = entry.get("size") + digest = entry.get("sha256") + if isinstance(size, bool) or not isinstance(size, int) or size < 0: + raise ValueError(f"invalid size for {relative}") + if not ( + isinstance(digest, str) + and len(digest) == 64 + and all(character in "0123456789abcdef" for character in digest) + ): + raise ValueError(f"invalid sha256 for {relative}") + if not saw_mcap: + raise ValueError("manifest contains no MCAP file") + return payload + + +def validate_episode_dir(directory: Path, manifest: dict[str, Any]) -> None: + ready = directory / "READY" + if not ready.is_file() or ready.is_symlink(): + raise ValueError("episode has no regular READY marker") + for entry in manifest["files"]: + relative = safe_relative_path(entry["path"]) + path = directory / relative + if not path.is_file() or path.is_symlink(): + raise ValueError(f"missing copied file: {relative}") + if path.stat().st_size != entry["size"]: + raise ValueError(f"size mismatch: {relative}") + if sha256_file(path) != entry["sha256"]: + raise ValueError(f"sha256 mismatch: {relative}") + + +class CommandRunner: + def run( + self, + command: Sequence[str], + *, + timeout: float, + ) -> subprocess.CompletedProcess[str]: + return subprocess.run( + list(command), + check=False, + capture_output=True, + text=True, + timeout=timeout, + ) + + +class DataGetSync: + def __init__( + self, + *, + remote: str, + remote_ready: str, + destination: Path, + status_file: Path, + runner: CommandRunner | None = None, + ssh_timeout_s: float = 8.0, + reserve_bytes: int = 1_073_741_824, + verify_existing: bool = False, + ) -> None: + self.remote = remote + self.remote_ready = remote_ready.rstrip("/") + self.destination = destination + self.status_file = status_file + self.runner = runner or CommandRunner() + self.ssh_timeout_s = ssh_timeout_s + self.reserve_bytes = reserve_bytes + self.verify_existing = verify_existing + self.stop_requested = False + self.started_at = time.time() + self.sync_count = 0 + self.last_episode: str | None = None + self.last_error = "" + + @property + def ssh_base(self) -> list[str]: + return [ + "ssh", + "-o", + "BatchMode=yes", + "-o", + f"ConnectTimeout={max(1, int(self.ssh_timeout_s))}", + self.remote, + ] + + def request_stop(self, _signum: int, _frame: object) -> None: + self.stop_requested = True + + def _write_status(self, state: str) -> None: + atomic_write_json( + self.status_file, + { + "state": state, + "remote": self.remote, + "remote_ready": self.remote_ready, + "destination": str(self.destination), + "sync_count": self.sync_count, + "last_episode": self.last_episode, + "last_error": self.last_error, + "uptime_s": round(time.time() - self.started_at, 1), + "updated_unix_s": time.time(), + }, + ) + + def _remote_command(self, command: str) -> str: + result = self.runner.run( + [*self.ssh_base, command], timeout=self.ssh_timeout_s + 2.0 + ) + if result.returncode != 0: + detail = result.stderr.strip() or result.stdout.strip() + raise RuntimeError(f"remote command failed: {detail}") + return result.stdout + + def list_remote_episodes(self) -> list[str]: + root = shlex.quote(self.remote_ready) + output = self._remote_command( + f"find {root} -mindepth 1 -maxdepth 1 -type d -printf '%f\\n'" + ) + names = sorted({line.strip() for line in output.splitlines() if line.strip()}) + unsafe = [name for name in names if not safe_episode_name(name)] + if unsafe: + raise RuntimeError(f"remote returned unsafe episode names: {unsafe!r}") + return names + + def get_remote_manifest(self, episode: str) -> dict[str, Any]: + if not safe_episode_name(episode): + raise ValueError(f"unsafe episode name: {episode!r}") + episode_path = f"{self.remote_ready}/{episode}" + manifest_path = f"{episode_path}/manifest.json" + ready_path = f"{episode_path}/READY" + raw = self._remote_command( + f"test -f {shlex.quote(ready_path)} && cat {shlex.quote(manifest_path)}" + ) + return parse_manifest(raw, episode) + + def _enough_local_space(self, manifest: dict[str, Any]) -> bool: + required = sum(int(entry["size"]) for entry in manifest["files"]) + free = shutil.disk_usage(self.destination).free + return free >= required + self.reserve_bytes + + def sync_episode(self, episode: str) -> bool: + if not safe_episode_name(episode): + raise ValueError(f"unsafe episode name: {episode!r}") + final = self.destination / episode + if final.exists(): + manifest_path = final / "manifest.json" + ready_path = final / "READY" + if ( + not manifest_path.is_file() + or manifest_path.is_symlink() + or not ready_path.is_file() + or ready_path.is_symlink() + ): + raise RuntimeError(f"existing destination is incomplete: {final}") + manifest = parse_manifest(manifest_path.read_text(encoding="utf-8"), episode) + # The first transfer already hashed every listed artifact before + # the atomic rename. Re-reading all historical MCAP files every + # two seconds would eventually saturate the workstation disk. + if self.verify_existing: + validate_episode_dir(final, manifest) + return False + + manifest = self.get_remote_manifest(episode) + if not self._enough_local_space(manifest): + raise RuntimeError("not enough local disk space for episode") + + incoming = self.destination / ".incoming" / episode + incoming.mkdir(parents=True, exist_ok=True) + remote_source = f"{self.remote}:{self.remote_ready}/{episode}/" + ssh_transport = ( + "ssh -o BatchMode=yes " + f"-o ConnectTimeout={max(1, int(self.ssh_timeout_s))}" + ) + result = self.runner.run( + [ + "rsync", + "-a", + "--partial", + "--protect-args", + "-e", + ssh_transport, + remote_source, + str(incoming) + "/", + ], + timeout=24 * 60 * 60, + ) + if result.returncode != 0: + detail = result.stderr.strip() or result.stdout.strip() + raise RuntimeError(f"rsync failed: {detail}") + + copied_manifest_path = incoming / "manifest.json" + if not copied_manifest_path.is_file(): + raise RuntimeError("copied episode has no manifest.json") + copied_manifest = parse_manifest( + copied_manifest_path.read_text(encoding="utf-8"), episode + ) + if copied_manifest != manifest: + raise RuntimeError("remote manifest changed during transfer") + validate_episode_dir(incoming, copied_manifest) + os.replace(incoming, final) + self.sync_count += 1 + self.last_episode = episode + return True + + def run_once(self) -> int: + self.destination.mkdir(parents=True, exist_ok=True) + (self.destination / ".incoming").mkdir(parents=True, exist_ok=True) + copied = 0 + errors: list[str] = [] + try: + episodes = self.list_remote_episodes() + for episode in episodes: + try: + if self.sync_episode(episode): + copied += 1 + except Exception as exc: + # One damaged historical episode must not starve newer + # ready data. Keep its error visible and continue. + errors.append(f"{episode}: {exc}") + if errors: + raise RuntimeError("; ".join(errors)) + self.last_error = "" + self._write_status("idle") + except Exception as exc: + self.last_error = str(exc) + self._write_status("error") + raise + return copied + + def run_forever(self, poll_seconds: float) -> None: + self.destination.mkdir(parents=True, exist_ok=True) + while not self.stop_requested: + try: + self.run_once() + except Exception: + pass + deadline = time.monotonic() + poll_seconds + while not self.stop_requested and time.monotonic() < deadline: + time.sleep(min(0.2, max(0.0, deadline - time.monotonic()))) + self._write_status("stopped") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--remote", default="nvidia@192.168.41.2") + parser.add_argument( + "--remote-ready", default="/home/nvidia/tg3_data_collection/ready" + ) + parser.add_argument( + "--destination", + type=Path, + default=Path( + "/home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/Data_Get" + ), + ) + parser.add_argument("--status-file", type=Path) + parser.add_argument("--poll-seconds", type=float, default=2.0) + parser.add_argument( + "--verify-existing", + action="store_true", + help="rehash already published local episodes (slow; intended for audits)", + ) + parser.add_argument("--once", action="store_true") + return parser.parse_args() + + +def main() -> int: + args = parse_args() + if args.poll_seconds <= 0.0: + raise SystemExit("--poll-seconds must be positive") + destination = args.destination.resolve() + status_file = args.status_file or destination / "sync_status.json" + syncer = DataGetSync( + remote=args.remote, + remote_ready=args.remote_ready, + destination=destination, + status_file=status_file, + verify_existing=args.verify_existing, + ) + signal.signal(signal.SIGINT, syncer.request_stop) + signal.signal(signal.SIGTERM, syncer.request_stop) + if args.once: + try: + syncer.run_once() + except Exception as exc: + print(f"data sync failed: {exc}") + return 1 + return 0 + syncer.run_forever(args.poll_seconds) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tg3_data_collection/test_data_get_sync.py b/tg3_data_collection/test_data_get_sync.py new file mode 100755 index 0000000..a3c3a70 --- /dev/null +++ b/tg3_data_collection/test_data_get_sync.py @@ -0,0 +1,149 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import hashlib +import json +import tempfile +import unittest +from pathlib import Path + +from data_get_sync import ( + DataGetSync, + parse_manifest, + safe_episode_name, + safe_relative_path, + validate_episode_dir, +) + + +class DataGetSyncTests(unittest.TestCase): + def test_safe_episode_name(self) -> None: + self.assertTrue(safe_episode_name("episode_20260810T120000000_deadbeef")) + self.assertFalse(safe_episode_name("../escape")) + self.assertFalse(safe_episode_name("bad/name")) + self.assertFalse(safe_episode_name("")) + + def test_safe_relative_path(self) -> None: + self.assertEqual(safe_relative_path("bag/metadata.yaml"), Path("bag/metadata.yaml")) + for invalid in ("", "/etc/passwd", "../escape", "bag/../../escape", None): + with self.subTest(invalid=invalid): + with self.assertRaises(ValueError): + safe_relative_path(invalid) + + def _episode(self, root: Path, name: str) -> tuple[Path, dict]: + episode = root / name + (episode / "bag").mkdir(parents=True) + mcap = episode / "bag" / "bag_0.mcap" + metadata = episode / "bag" / "metadata.yaml" + mcap.write_bytes(b"mcap-data") + metadata.write_text("rosbag2_bagfile_information: {}\n", encoding="utf-8") + files = [] + for path in (mcap, metadata): + raw = path.read_bytes() + files.append( + { + "path": str(path.relative_to(episode)), + "size": len(raw), + "sha256": hashlib.sha256(raw).hexdigest(), + } + ) + manifest = {"state": "complete", "episode_id": name, "files": files} + (episode / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (episode / "READY").write_text("ready\n", encoding="ascii") + return episode, manifest + + def test_parse_and_validate_episode(self) -> None: + with tempfile.TemporaryDirectory() as temporary: + name = "episode_20260810T120000000_deadbeef" + episode, manifest = self._episode(Path(temporary), name) + parsed = parse_manifest(json.dumps(manifest), name) + validate_episode_dir(episode, parsed) + + def test_manifest_requires_complete_matching_episode_and_mcap(self) -> None: + digest = "0" * 64 + base = { + "state": "complete", + "episode_id": "episode_ok", + "files": [{"path": "bag/bag_0.mcap", "size": 1, "sha256": digest}], + } + with self.assertRaises(ValueError): + parse_manifest(json.dumps({**base, "state": "active"}), "episode_ok") + with self.assertRaises(ValueError): + parse_manifest(json.dumps(base), "episode_other") + no_mcap = { + **base, + "files": [{"path": "bag/metadata.yaml", "size": 1, "sha256": digest}], + } + with self.assertRaises(ValueError): + parse_manifest(json.dumps(no_mcap), "episode_ok") + + def test_validate_detects_tampering(self) -> None: + with tempfile.TemporaryDirectory() as temporary: + name = "episode_20260810T120000000_deadbeef" + episode, manifest = self._episode(Path(temporary), name) + (episode / "bag" / "bag_0.mcap").write_bytes(b"changed") + with self.assertRaises(ValueError): + validate_episode_dir(episode, manifest) + + def test_validate_requires_regular_ready_marker(self) -> None: + with tempfile.TemporaryDirectory() as temporary: + name = "episode_20260810T120000000_deadbeef" + episode, manifest = self._episode(Path(temporary), name) + (episode / "READY").unlink() + with self.assertRaises(ValueError): + validate_episode_dir(episode, manifest) + + def test_existing_episode_skips_expensive_hash_unless_requested(self) -> None: + with tempfile.TemporaryDirectory() as temporary: + root = Path(temporary) + name = "episode_20260810T120000000_deadbeef" + episode, _manifest = self._episode(root, name) + (episode / "bag" / "bag_0.mcap").write_bytes(b"changed") + fast = DataGetSync( + remote="unused", + remote_ready="/unused", + destination=root, + status_file=root / "status.json", + ) + self.assertFalse(fast.sync_episode(name)) + deep = DataGetSync( + remote="unused", + remote_ready="/unused", + destination=root, + status_file=root / "status.json", + verify_existing=True, + ) + with self.assertRaises(ValueError): + deep.sync_episode(name) + + def test_bad_old_episode_does_not_starve_newer_episode(self) -> None: + with tempfile.TemporaryDirectory() as temporary: + root = Path(temporary) + + class ProbeSync(DataGetSync): + def __init__(self) -> None: + super().__init__( + remote="unused", + remote_ready="/unused", + destination=root, + status_file=root / "status.json", + ) + self.seen: list[str] = [] + + def list_remote_episodes(self) -> list[str]: + return ["episode_bad", "episode_new"] + + def sync_episode(self, episode: str) -> bool: + self.seen.append(episode) + if episode == "episode_bad": + raise RuntimeError("damaged") + return True + + syncer = ProbeSync() + with self.assertRaises(RuntimeError): + syncer.run_once() + self.assertEqual(syncer.seen, ["episode_bad", "episode_new"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tg3_data_collection/tg3-data-get-sync.service b/tg3_data_collection/tg3-data-get-sync.service new file mode 100644 index 0000000..6c092d7 --- /dev/null +++ b/tg3_data_collection/tg3-data-get-sync.service @@ -0,0 +1,14 @@ +[Unit] +Description=TG3 completed data episode sync to Data_Get + +[Service] +Type=simple +WorkingDirectory=/home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop +ExecStart=/usr/bin/python3 /home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/tg3_data_collection/data_get_sync.py --remote nvidia@192.168.41.2 --remote-ready /home/nvidia/tg3_data_collection/ready --destination /home/ps/Desktop/TG3_TS1P_OmniSocket_Teleop/Data_Get --poll-seconds 2 +Restart=always +RestartSec=2 +KillSignal=SIGINT +TimeoutStopSec=10 + +[Install] +WantedBy=default.target diff --git a/tg3_local_teleop/README.md b/tg3_local_teleop/README.md index adb86d1..ac4a0d3 100644 --- a/tg3_local_teleop/README.md +++ b/tg3_local_teleop/README.md @@ -145,13 +145,70 @@ OmniSocket 进程;在此之前桥会保持最后一帧目标。未 START 或 正常待机,不触发反复重启。只有匹配的操作员 STOP 自动回 Home;意外断网不自动产生 回位运动。 +## 独立数采 + +数采使用左摇杆按下键(L3),不是左侧 X/Y/Z 面键: + +- 遥操已成功开启后,先保持 L3 松开至少 `0.5 s`; +- 连续按住 L3 `1 s` 开始一条 episode; +- 松开至少 `0.5 s`,再次连续按住 `1 s` 正常结束; +- 长按 Z+C 结束遥操、安全解除或桥退出时,也会异步请求结束当前 episode; +- 数采失败只记录错误,不解除遥操、不阻塞 50 Hz 控制,也不延迟 STOP/Home。 + +`button_joystick.left` 是 xTELE 5003 中独立的摇杆按压字段。已安装的 xTELE 0.1.2 +没有给它注册处理函数;X/Y/Z/A/B/C 均已有厂家功能或本项目绑定,因此不复用面键。 +新会话和服务重启后都先锁定为“必须松开”,畸形或陈旧按键帧不能被当作有效松开。 + +项目自有 `tg3-data-recorder.service` 在 Nvidia 上运行独立 `ros2 bag record`,不停止、 +重配或接管 Ubuntu 厂家 `/record_bag_node`。默认仅录 `config.toml` 中的明确白名单, +包含 `/robot_state`、双臂/BrainCo 双手命令与反馈、HBWALK、IMU、电源状态以及完整 +xTELE 应用帧;相机和点云默认不录,也禁止改成 `-a`。 + +机器人暂存目录: + +```text +/home/nvidia/tg3_data_collection/ + active/ # 尚未完成,不能取走 + ready/ # 已正常收尾、校验并生成 READY + failed/ # 启动、磁盘、进程或校验失败,保留供诊断 +``` + +MCAP 使用 `zstd_fast`、64 MiB cache、每 300 秒分片;每条 episode 最长 30 分钟, +启动/运行最低保留 20 GiB。正常停止以 SIGINT 让 rosbag 写完 `metadata.yaml`,随后执行 +`ros2 bag info`,并确认所有 required topics 都有非零消息数;然后计算 SHA-256、写 +`manifest.json` 和 `READY`,最后才原子进入 `ready/`。Nvidia 需已安装 `python3-yaml` +(现场已验证 PyYAML 6.0.1)。 + +部署两个机器人侧服务: + +```bash +mkdir -p ~/.config/systemd/user +cp tg3-data-recorder.service tg3-local-teleop.service ~/.config/systemd/user/ +systemctl --user daemon-reload +systemctl --user enable --now tg3-data-recorder.service +systemctl --user restart tg3-local-teleop.service +``` + +只读检查: + +```bash +systemctl --user --no-pager status tg3-data-recorder.service +ros2 topic echo --once /tg3/data_collection/status +python3 -m json.tool /home/nvidia/tg3_local_teleop/status.json +find /home/nvidia/tg3_data_collection/ready -mindepth 1 -maxdepth 1 -type d +``` + +完成 episode 由 PS 本机服务校验后保存到项目 `Data_Get`,详见 +`../tg3_data_collection/README.md`。机器人 `ready/` 中的副本不会自动删除。 + ## 完整重启顺序 只重启本项目的网络与机器人桥时,先启动 Nvidia 接收端,再启动 EAI 本地门控服务; EAI 在收到物理 START 前不会建立发送 Session: ```bash -ssh nvidia@192.168.41.2 'systemctl --user restart tg3-local-teleop.service' +ssh nvidia@192.168.41.2 \ + 'systemctl --user restart tg3-data-recorder.service tg3-local-teleop.service' ssh eai 'systemctl --user restart tg3-omnisocket-sender.service' ``` diff --git a/tg3_local_teleop/config.toml b/tg3_local_teleop/config.toml index 5c725f2..f2adde7 100644 --- a/tg3_local_teleop/config.toml +++ b/tg3_local_teleop/config.toml @@ -1,7 +1,7 @@ [network] # Direct OmniSocket KCP input. No robot-side ZMQ receiving proxy is used. transport = "omnisocket" -omnisocket_server = "175.178.116.187:14049" +omnisocket_server = "192.168.5.14:14049" omnisocket_peer_id = "tg3-009027fa8190-robot" omnisocket_expected_sender = "tg3-009027fa8190-iarm" omnisocket_max_packet_age_ms = 300.0 @@ -55,6 +55,71 @@ joint_upper_rad = [ 2.8449, 0.1920, 2.8449, 0.1920, 2.8449, 1.3265, 1.3265, ] +[data_collection] +# Left joystick press (L3) is a dedicated raw xTELE field and is not registered +# by the installed xTELE 0.1.2 button handlers. It only toggles recording while +# a validated teleoperation session is armed. Z+C STOP/safety disarm always +# ends an active recording without delaying robot control or Home. +enabled = true +button_hold_seconds = 1.0 +button_release_seconds = 0.5 +# This threshold qualifies only the L3 button edge; it is not a robot-motion +# watchdog and does not restore the removed 0.25 s network disarm gate. +button_input_timeout_s = 0.25 +control_retry_seconds = 0.5 +ack_timeout_seconds = 5.0 +heartbeat_interval_seconds = 0.5 +heartbeat_timeout_seconds = 3.0 +status_stale_seconds = 4.0 +control_topic = "/tg3/data_collection/control" +status_topic = "/tg3/data_collection/status" +iarm_frame_topic = "/tg3/data_collection/iarm_frame" +base_directory = "/home/nvidia/tg3_data_collection" +minimum_free_gib = 20.0 +max_duration_seconds = 1800.0 + +# Explicit control-data whitelist. Camera images and point clouds are omitted +# by default to bound disk and network load; never replace this with `-a` on the +# live robot. /robot_state is the authoritative measured robot state. +topics = [ + "/robot_state", + "/encoder_identical_joint", + "/arm/cmd", + "/freq_change/arm_status", + "/data_logger/arm_status", + "/left_hand/set_motor_multi", + "/right_hand/set_motor_multi", + "/left_hand/motor_status", + "/right_hand/motor_status", + "/left_hand/touch_status", + "/right_hand/touch_status", + "/hric/robot/cmd_vel", + "/hric/robot/cmd_vel_status", + "/hric/robot/rl_state", + "/imu_data", + "/power/board/key_status", + "/power/board/status", + "/power/battery/status", + "/tg3/data_collection/control", + "/tg3/data_collection/status", + "/tg3/data_collection/iarm_frame", +] + +# Refuse a new episode if these core streams are absent from the live ROS +# graph. Other whitelisted streams may appear later and rosbag discovery will +# subscribe to them without restarting the recorder. +required_topics = [ + "/robot_state", + "/encoder_identical_joint", + "/freq_change/arm_status", + "/left_hand/motor_status", + "/right_hand/motor_status", + "/hric/robot/rl_state", + "/tg3/data_collection/control", + "/tg3/data_collection/status", + "/tg3/data_collection/iarm_frame", +] + [locomotion] # The two original immediate bindings are independent: right C + left-stick # vertical controls translation; left Z + right-stick horizontal controls diff --git a/tg3_local_teleop/data_collection.py b/tg3_local_teleop/data_collection.py new file mode 100755 index 0000000..43bcc67 --- /dev/null +++ b/tg3_local_teleop/data_collection.py @@ -0,0 +1,1140 @@ +#!/usr/bin/env python3 +"""Session-safe, asynchronous ROS 2 bag data collection helpers. + +The module intentionally has no ROS Python dependency. The teleoperation +bridge can call the small, non-blocking API from its control tick while a +worker thread owns all filesystem work and the ``ros2 bag`` subprocess. +""" + +from __future__ import annotations + +import hashlib +import json +import math +import os +import queue +import shutil +import signal +import subprocess +import threading +import time +import uuid +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Literal, Protocol + +import yaml + + +ToggleAction = Literal["start", "stop"] + + +def left_joystick_pressed(data: Mapping[str, Any]) -> bool | None: + """Strictly parse the live-frame ``button_joystick.left`` value. + + xTELE's live 5003 frame carries this button directly as ``bool`` or the + integer ``0``/``1``. Missing containers and all other values return + ``None``. In particular, callers must not interpret malformed input as a + button release. + """ + + if not isinstance(data, Mapping): + return None + try: + button_joystick = data["button_joystick"] + if not isinstance(button_joystick, Mapping): + return None + value = button_joystick["left"] + except (KeyError, TypeError): + return None + if type(value) is bool: + return value + if type(value) is int and value in (0, 1): + return bool(value) + return None + + +class RecordingToggleGate: + """One-button start/stop gate with long-press and release protection. + + Each accepted teleoperation session must call :meth:`new_session`. The + gate then requires a continuous, valid release before accepting its first + hold. A hold toggles once; the same physical press can never toggle a + second time. Missing/malformed input cancels timers and re-enters the + release lock. + """ + + def __init__( + self, hold_seconds: float = 1.0, release_seconds: float = 0.5 + ) -> None: + if not math.isfinite(hold_seconds) or hold_seconds <= 0.0: + raise ValueError("recording hold time must be positive and finite") + if not math.isfinite(release_seconds) or release_seconds <= 0.0: + raise ValueError("recording release time must be positive and finite") + self.hold_seconds = float(hold_seconds) + self.release_seconds = float(release_seconds) + self.session_open = False + self.active = False + self.require_release = True + self.hold_started_at: float | None = None + self.release_started_at: float | None = None + self.input_healthy = False + self.button_pressed: bool | None = None + self.toggle_count = 0 + self.last_transition = "initialized" + + def new_session(self) -> None: + """Open a fresh session in the release-locked, inactive state.""" + + self.session_open = True + self.active = False + self.require_release = True + self.hold_started_at = None + self.release_started_at = None + self.input_healthy = False + self.button_pressed = None + self.last_transition = "new_session" + + def end_session(self, reason: str = "session_ended") -> ToggleAction | None: + """Close the gate and report whether an active recorder must stop.""" + + action: ToggleAction | None = "stop" if self.active else None + self.session_open = False + self.active = False + self.require_release = True + self.hold_started_at = None + self.release_started_at = None + self.input_healthy = False + self.button_pressed = None + self.last_transition = reason + return action + + def force_inactive(self, reason: str = "recorder_stopped") -> None: + """Reconcile the gate after an automatic or failed recorder stop.""" + + self.active = False + self.require_release = True + self.hold_started_at = None + self.release_started_at = None + self.input_healthy = False + self.button_pressed = None + self.last_transition = reason + + def update( + self, + now: float, + *, + input_healthy: bool, + pressed: bool | None, + ) -> ToggleAction | None: + """Advance the gate and return ``start``/``stop`` only on a toggle.""" + + if not math.isfinite(now): + raise ValueError("recording gate clock must be finite") + if not self.session_open: + return None + + self.input_healthy = bool(input_healthy) + self.button_pressed = pressed if type(pressed) is bool else None + if not input_healthy or pressed is None or type(pressed) is not bool: + self.hold_started_at = None + self.release_started_at = None + self.require_release = True + return None + + if self.require_release: + self.hold_started_at = None + if pressed: + self.release_started_at = None + return None + if self.release_started_at is None: + self.release_started_at = now + return None + if now - self.release_started_at >= self.release_seconds: + self.require_release = False + self.release_started_at = None + return None + + self.release_started_at = None + if not pressed: + if self.hold_started_at is not None: + # A short press must be followed by another stable release. + self.hold_started_at = None + self.release_started_at = now + self.require_release = True + return None + + if self.hold_started_at is None: + self.hold_started_at = now + return None + if now - self.hold_started_at < self.hold_seconds: + return None + + self.active = not self.active + self.toggle_count += 1 + action: ToggleAction = "start" if self.active else "stop" + self.last_transition = action + self.hold_started_at = None + self.release_started_at = None + self.require_release = True + return action + + @property + def state(self) -> str: + if not self.session_open: + return "closed" + if not self.input_healthy or self.button_pressed is None: + return "input_unhealthy" + if self.require_release: + return "awaiting_release" + if self.hold_started_at is not None: + return "holding_stop" if self.active else "holding_start" + return "active" if self.active else "idle" + + +@dataclass(frozen=True) +class RecorderConfig: + """Static configuration for :class:`DataRecorderManager`.""" + + base_directory: Path | str = Path("/home/nvidia/tg3_data_collection") + topics: Sequence[str] = () + required_topics: Sequence[str] = () + minimum_free_bytes: int = 5 * 1024**3 + max_duration_seconds: float = 30 * 60.0 + poll_interval_seconds: float = 0.1 + sigint_timeout_seconds: float = 15.0 + kill_timeout_seconds: float = 3.0 + ros2_executable: str = "ros2" + validate_bag_info: bool = True + bag_info_timeout_seconds: float = 15.0 + + def __post_init__(self) -> None: + base = Path(self.base_directory).expanduser() + topics = tuple(self.topics) + required_topics = tuple(self.required_topics) + if not topics: + raise ValueError("at least one recording topic is required") + if any( + not isinstance(topic, str) + or not topic.startswith("/") + or topic.strip() != topic + or any(character.isspace() for character in topic) + for topic in topics + ): + raise ValueError("recording topics must be absolute ROS topic names") + if len(set(topics)) != len(topics): + raise ValueError("recording topics must not contain duplicates") + if len(set(required_topics)) != len(required_topics): + raise ValueError("required topics must not contain duplicates") + if any( + not isinstance(topic, str) + or not topic.startswith("/") + or topic.strip() != topic + or any(character.isspace() for character in topic) + for topic in required_topics + ): + raise ValueError("required topics must be absolute ROS topic names") + if not set(required_topics).issubset(topics): + raise ValueError("required topics must be a subset of recording topics") + if ( + isinstance(self.minimum_free_bytes, bool) + or not isinstance(self.minimum_free_bytes, int) + or self.minimum_free_bytes < 0 + ): + raise ValueError("minimum free bytes must be a non-negative integer") + for name, value in ( + ("maximum duration", self.max_duration_seconds), + ("poll interval", self.poll_interval_seconds), + ("SIGINT timeout", self.sigint_timeout_seconds), + ("kill timeout", self.kill_timeout_seconds), + ("bag info timeout", self.bag_info_timeout_seconds), + ): + if not math.isfinite(value) or value <= 0.0: + raise ValueError(f"{name} must be positive and finite") + if not isinstance(self.ros2_executable, str) or not self.ros2_executable: + raise ValueError("ROS 2 executable must be a non-empty string") + if type(self.validate_bag_info) is not bool: + raise ValueError("validate bag info must be a bool") + object.__setattr__(self, "base_directory", base) + object.__setattr__(self, "topics", topics) + object.__setattr__(self, "required_topics", required_topics) + + +class RecorderProcess(Protocol): + returncode: int | None + + def poll(self) -> int | None: ... + + def send_signal(self, sig: int) -> None: ... + + def wait(self, timeout: float | None = None) -> int: ... + + def terminate(self) -> None: ... + + def kill(self) -> None: ... + + +ProcessFactory = Callable[..., RecorderProcess] + + +@dataclass(frozen=True) +class _StartRequest: + episode_id: str + session_id: str + reason: str + + +@dataclass +class _Episode: + request: _StartRequest + active_directory: Path + bag_directory: Path + started_monotonic: float + started_at: datetime + node_name: str = "" + command: tuple[str, ...] = () + bag_info_validation: dict[str, Any] | None = None + required_topic_message_counts: dict[str, int] | None = None + process: RecorderProcess | None = None + stdout_stream: Any = None + stderr_stream: Any = None + + +class RecordingError(RuntimeError): + """An episode could not be safely finalized as usable data.""" + + +class DataRecorderManager: + """Asynchronous owner of one ``ros2 bag record`` process at a time. + + :meth:`request_start` and :meth:`request_stop` only update in-memory state + and enqueue work. The worker performs every subprocess and filesystem + operation. Completed episodes move atomically from ``active`` to + ``ready``; any startup, process, validation, or finalization failure moves + the preserved episode directory to ``failed``. + """ + + def __init__( + self, + config: RecorderConfig, + *, + process_factory: ProcessFactory = subprocess.Popen, + command_runner: Callable[..., Any] = subprocess.run, + free_bytes: Callable[[Path], int] | None = None, + monotonic: Callable[[], float] = time.monotonic, + utc_now: Callable[[], datetime] | None = None, + episode_id_factory: Callable[[], str] | None = None, + ) -> None: + self.config = config + self._process_factory = process_factory + self._command_runner = command_runner + self._free_bytes = free_bytes or ( + lambda path: int(shutil.disk_usage(path).free) + ) + self._monotonic = monotonic + self._utc_now = utc_now or (lambda: datetime.now(timezone.utc)) + self._episode_id_factory = episode_id_factory or self._default_episode_id + self._requests: queue.Queue[tuple[str, object | None]] = queue.Queue() + self._condition = threading.Condition(threading.RLock()) + self._shutdown_requested = False + self._requested_stop_reason: str | None = None + self._status: dict[str, Any] = { + "state": "idle", + "recording": False, + "episode_id": None, + "session_id": None, + "started_at": None, + "stop_reason": None, + "last_result": None, + "last_episode_directory": None, + "last_error": "", + "topics": list(config.topics), + } + self._worker = threading.Thread( + target=self._run, + name="tg3-data-recorder", + daemon=True, + ) + self._worker.start() + + @staticmethod + def _default_episode_id() -> str: + stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S.%fZ") + return f"{stamp}_{uuid.uuid4().hex[:12]}" + + def request_start( + self, + session_id: str, + reason: str = "button", + *, + episode_id: str | None = None, + ) -> str | None: + """Queue a start and return its actual episode ID, or ``None`` if busy.""" + + self._validate_custom_data_value(session_id, "session ID") + self._validate_text(reason, "start reason") + if episode_id is None: + episode_id = self._episode_id_factory() + self._validate_episode_id(episode_id) + request = _StartRequest(episode_id, session_id, reason) + with self._condition: + if self._shutdown_requested or self._status["state"] != "idle": + return None + self._requested_stop_reason = None + self._status.update( + { + "state": "start_pending", + "recording": False, + "episode_id": episode_id, + "session_id": session_id, + "started_at": None, + "stop_reason": None, + "last_error": "", + } + ) + self._condition.notify_all() + self._requests.put(("start", request)) + return episode_id + + def request_stop(self, reason: str = "operator_button") -> bool: + """Queue SIGINT shutdown of the active/pending episode.""" + + self._validate_text(reason, "stop reason") + with self._condition: + if self._status["state"] not in ( + "start_pending", + "starting", + "recording", + "stop_pending", + ): + return False + if self._requested_stop_reason is None: + self._requested_stop_reason = reason + self._status["state"] = "stop_pending" + self._status["stop_reason"] = self._requested_stop_reason + self._condition.notify_all() + self._requests.put(("stop", reason)) + return True + + def status(self) -> dict[str, Any]: + """Return a thread-safe, JSON-serializable status snapshot.""" + + with self._condition: + snapshot = dict(self._status) + snapshot["topics"] = list(self._status["topics"]) + return snapshot + + def wait_until_idle(self, timeout: float | None = None) -> bool: + """Wait explicitly for an idle/terminated worker (never call in tick).""" + + deadline = None if timeout is None else self._monotonic() + timeout + with self._condition: + while self._status["state"] not in ("idle", "shutdown"): + remaining = ( + None if deadline is None else deadline - self._monotonic() + ) + if remaining is not None and remaining <= 0.0: + return False + self._condition.wait(remaining) + return True + + def shutdown(self, *, wait: bool = True, timeout: float | None = None) -> bool: + """Request clean SIGINT shutdown and optionally wait for the worker.""" + + with self._condition: + if not self._shutdown_requested: + self._shutdown_requested = True + if self._requested_stop_reason is None: + self._requested_stop_reason = "shutdown" + if self._status["state"] not in ("idle", "shutdown"): + self._status["state"] = "stop_pending" + self._status["stop_reason"] = self._requested_stop_reason + self._condition.notify_all() + self._requests.put(("shutdown", None)) + if not wait: + return not self._worker.is_alive() + self._worker.join(timeout) + return not self._worker.is_alive() + + def _run(self) -> None: + while True: + command, payload = self._requests.get() + if command == "shutdown": + break + if command != "start" or not isinstance(payload, _StartRequest): + continue + self._record_episode(payload) + with self._condition: + should_shutdown = self._shutdown_requested + if should_shutdown: + break + with self._condition: + self._status["state"] = "shutdown" + self._status["recording"] = False + self._status["episode_id"] = None + self._status["session_id"] = None + self._condition.notify_all() + + def _record_episode(self, request: _StartRequest) -> None: + episode: _Episode | None = None + stop_reason: str | None = None + result_directory: Path | None = None + try: + with self._condition: + self._status["state"] = "starting" + self._condition.notify_all() + + base = Path(self.config.base_directory) + active_root = base / "active" + ready_root = base / "ready" + failed_root = base / "failed" + for directory in (active_root, ready_root, failed_root): + directory.mkdir(parents=True, exist_ok=True) + + active_directory = active_root / request.episode_id + active_directory.mkdir(mode=0o750) + episode = _Episode( + request=request, + active_directory=active_directory, + bag_directory=active_directory / "bag", + started_monotonic=self._monotonic(), + started_at=self._as_utc(self._utc_now()), + bag_info_validation={ + "enabled": self.config.validate_bag_info, + "passed": None, + "result": "not_run", + "timeout_seconds": self.config.bag_info_timeout_seconds, + }, + ) + initial_free = int(self._free_bytes(base)) + if initial_free < self.config.minimum_free_bytes: + raise RecordingError( + "insufficient free space before recording: " + f"{initial_free} < {self.config.minimum_free_bytes} bytes" + ) + + episode.stdout_stream = (active_directory / "ros2_bag.stdout.log").open( + "ab", buffering=0 + ) + episode.stderr_stream = (active_directory / "ros2_bag.stderr.log").open( + "ab", buffering=0 + ) + command = [ + self.config.ros2_executable, + "bag", + "record", + "--storage", + "mcap", + "--storage-preset-profile", + "zstd_fast", + "--max-cache-size", + "67108864", + "--max-bag-duration", + "300", + "--disable-keyboard-controls", + "--node-name", + self._ros_node_name(request.episode_id), + "--output", + str(episode.bag_directory), + "--custom-data", + f"capture_id={request.episode_id}", + f"teleop_session_id={request.session_id}", + "--topics", + *self.config.topics, + ] + episode.node_name = command[command.index("--node-name") + 1] + episode.command = tuple(command) + episode.process = self._process_factory( + command, + stdin=subprocess.DEVNULL, + stdout=episode.stdout_stream, + stderr=episode.stderr_stream, + start_new_session=True, + ) + with self._condition: + self._status.update( + { + "state": ( + "stop_pending" + if self._requested_stop_reason is not None + else "recording" + ), + "recording": True, + "started_at": self._iso(episode.started_at), + } + ) + self._condition.notify_all() + + stop_reason = self._monitor_episode(episode) + self._stop_process(episode.process) + ended_at = self._as_utc(self._utc_now()) + ended_monotonic = self._monotonic() + self._close_logs(episode) + with self._condition: + self._status["state"] = "finalizing" + self._status["recording"] = False + self._status["stop_reason"] = stop_reason + self._condition.notify_all() + + manifest = self._ready_manifest( + episode, + stop_reason, + ended_at, + ended_monotonic, + ) + self._write_manifest(episode.active_directory, manifest) + self._write_ready_marker(episode.active_directory) + self._sync_directory(episode.active_directory) + result_directory = ready_root / request.episode_id + os.replace(episode.active_directory, result_directory) + try: + self._sync_directory(ready_root) + except OSError: + # The same-filesystem rename has already published a complete + # manifest+READY directory. Parent fsync is best-effort here. + pass + self._finish_status("ready", result_directory, "", stop_reason) + except Exception as error: + error_text = f"{type(error).__name__}: {error}" + if episode is not None: + self._abort_process(episode.process) + self._close_logs(episode) + result_directory, failure_error = self._preserve_failed_episode( + episode, + stop_reason or self._requested_stop_reason or "error", + error_text, + ) + if failure_error: + error_text += f"; failed to archive episode: {failure_error}" + self._finish_status( + "failed", + result_directory, + error_text, + stop_reason or self._requested_stop_reason or "error", + ) + + def _monitor_episode(self, episode: _Episode) -> str: + assert episode.process is not None + while True: + returncode = episode.process.poll() + if returncode is not None: + raise RecordingError( + f"ros2 bag exited unexpectedly with return code {returncode}" + ) + + with self._condition: + requested_reason = self._requested_stop_reason + shutting_down = self._shutdown_requested + if requested_reason is not None: + return requested_reason + if shutting_down: + return "shutdown" + + elapsed = self._monotonic() - episode.started_monotonic + if elapsed >= self.config.max_duration_seconds: + return "maximum_duration" + free = int(self._free_bytes(Path(self.config.base_directory))) + if free < self.config.minimum_free_bytes: + return "low_disk_space" + + try: + command, payload = self._requests.get( + timeout=self.config.poll_interval_seconds + ) + except queue.Empty: + continue + if command == "stop": + return str(payload) + if command == "shutdown": + return "shutdown" + + def _stop_process(self, process: RecorderProcess) -> None: + if process.poll() is not None: + return + process.send_signal(signal.SIGINT) + try: + process.wait(timeout=self.config.sigint_timeout_seconds) + return + except subprocess.TimeoutExpired: + process.terminate() + try: + process.wait(timeout=self.config.kill_timeout_seconds) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=self.config.kill_timeout_seconds) + raise RecordingError("ros2 bag did not stop cleanly after SIGINT") + + def _abort_process(self, process: RecorderProcess | None) -> None: + if process is None or process.poll() is not None: + return + try: + process.send_signal(signal.SIGINT) + process.wait(timeout=self.config.sigint_timeout_seconds) + return + except Exception: + pass + try: + process.kill() + process.wait(timeout=self.config.kill_timeout_seconds) + except Exception: + pass + + def _ready_manifest( + self, + episode: _Episode, + stop_reason: str, + ended_at: datetime, + ended_monotonic: float, + ) -> dict[str, Any]: + metadata = episode.bag_directory / "metadata.yaml" + if not metadata.is_file() or metadata.stat().st_size <= 0: + raise RecordingError("bag metadata.yaml is missing or empty") + try: + metadata_text = metadata.read_text(encoding="utf-8") + except UnicodeDecodeError as error: + raise RecordingError("bag metadata.yaml is not valid UTF-8") from error + if "mcap" not in metadata_text.lower(): + raise RecordingError("bag metadata.yaml does not identify MCAP storage") + required_topic_counts = self._validate_required_topic_counts( + episode, metadata_text + ) + + mcap_files = sorted( + path + for path in episode.bag_directory.rglob("*.mcap") + if path.is_file() + ) + if not mcap_files: + raise RecordingError("bag contains no MCAP file") + empty_mcap = [path.name for path in mcap_files if path.stat().st_size <= 0] + if empty_mcap: + raise RecordingError( + "bag contains empty MCAP file(s): " + ", ".join(empty_mcap) + ) + bag_info_validation = self._validate_bag_info(episode) + artifacts = [metadata, *mcap_files] + files = [ + self._file_record(episode.active_directory, path) + for path in artifacts + ] + return { + "schema_version": 1, + "state": "complete", + "status": "ready", + "episode_id": episode.request.episode_id, + "session_id": episode.request.session_id, + "start_reason": episode.request.reason, + "stop_reason": stop_reason, + "topics": list(self.config.topics), + "required_topics": list(self.config.required_topics), + "required_topic_message_counts": required_topic_counts, + "storage_id": "mcap", + "custom_data": { + "capture_id": episode.request.episode_id, + "teleop_session_id": episode.request.session_id, + }, + "bag_info_validation": bag_info_validation, + "rosbag2": { + "node_name": episode.node_name, + "storage_preset_profile": "zstd_fast", + "max_cache_size_bytes": 67_108_864, + "max_bag_duration_seconds": 300, + }, + "started_at_utc": self._iso(episode.started_at), + "stopped_at_utc": self._iso(ended_at), + "duration_seconds": max( + 0.0, ended_monotonic - episode.started_monotonic + ), + "total_size_bytes": sum(record["size_bytes"] for record in files), + "files": files, + } + + def _validate_required_topic_counts( + self, episode: _Episode, metadata_text: str + ) -> dict[str, int]: + """Require positive rosbag2 metadata counts for configured topics.""" + + try: + metadata = yaml.safe_load(metadata_text) + except yaml.YAMLError as error: + raise RecordingError("bag metadata.yaml is not valid YAML") from error + if not isinstance(metadata, Mapping): + raise RecordingError("bag metadata.yaml root must be a mapping") + information = metadata.get("rosbag2_bagfile_information") + if not isinstance(information, Mapping): + raise RecordingError( + "bag metadata.yaml lacks rosbag2_bagfile_information" + ) + if information.get("storage_identifier") != "mcap": + raise RecordingError( + "bag metadata.yaml storage_identifier is not mcap" + ) + + entries = information.get("topics_with_message_count") + if not isinstance(entries, list): + if self.config.required_topics: + raise RecordingError( + "bag metadata.yaml lacks topics_with_message_count" + ) + episode.required_topic_message_counts = {} + return {} + + all_counts: dict[str, int] = {} + for index, entry in enumerate(entries): + if not isinstance(entry, Mapping): + raise RecordingError( + f"metadata topic-count entry {index} is not a mapping" + ) + topic_metadata = entry.get("topic_metadata") + if not isinstance(topic_metadata, Mapping): + raise RecordingError( + f"metadata topic-count entry {index} lacks topic_metadata" + ) + name = topic_metadata.get("name") + message_count = entry.get("message_count") + if not isinstance(name, str) or not name.startswith("/"): + raise RecordingError( + f"metadata topic-count entry {index} has an invalid name" + ) + if ( + isinstance(message_count, bool) + or not isinstance(message_count, int) + or message_count < 0 + ): + raise RecordingError( + f"metadata topic {name!r} has an invalid message_count" + ) + all_counts[name] = all_counts.get(name, 0) + message_count + + required_counts = { + topic: all_counts.get(topic, 0) + for topic in self.config.required_topics + } + episode.required_topic_message_counts = required_counts + missing = [ + topic for topic in self.config.required_topics if topic not in all_counts + ] + empty = [ + topic + for topic in self.config.required_topics + if topic in all_counts and all_counts[topic] <= 0 + ] + problems: list[str] = [] + if missing: + problems.append("missing required topics: " + ", ".join(missing)) + if empty: + problems.append("zero-message required topics: " + ", ".join(empty)) + if problems: + raise RecordingError("; ".join(problems)) + return required_counts + + def _validate_bag_info(self, episode: _Episode) -> dict[str, Any]: + """Run rosbag2's reader-level validation before publishing READY.""" + + if not self.config.validate_bag_info: + validation = { + "enabled": False, + "passed": None, + "result": "disabled", + "timeout_seconds": self.config.bag_info_timeout_seconds, + } + episode.bag_info_validation = validation + return validation + + command = [ + self.config.ros2_executable, + "bag", + "info", + str(episode.bag_directory), + ] + validation: dict[str, Any] = { + "enabled": True, + "passed": False, + "result": "running", + "returncode": None, + "timeout_seconds": self.config.bag_info_timeout_seconds, + "validated_at_utc": self._iso(self._as_utc(self._utc_now())), + } + episode.bag_info_validation = validation + try: + result = self._command_runner( + command, + stdin=subprocess.DEVNULL, + capture_output=True, + text=True, + timeout=self.config.bag_info_timeout_seconds, + check=False, + ) + except subprocess.TimeoutExpired as error: + validation["result"] = "timeout" + self._write_bag_info_logs( + episode, + self._output_text(error.stdout), + self._output_text(error.stderr), + ) + raise RecordingError( + "ros2 bag info timed out after " + f"{self.config.bag_info_timeout_seconds:.3f}s" + ) from error + except Exception as error: + validation["result"] = "invocation_error" + validation["error_type"] = type(error).__name__ + raise RecordingError( + f"could not run ros2 bag info: {type(error).__name__}: {error}" + ) from error + + returncode = getattr(result, "returncode", None) + stdout = self._output_text(getattr(result, "stdout", "")) + stderr = self._output_text(getattr(result, "stderr", "")) + self._write_bag_info_logs(episode, stdout, stderr) + if type(returncode) is not int: + validation["result"] = "invalid_runner_result" + raise RecordingError("ros2 bag info returned no integer return code") + validation["returncode"] = returncode + if returncode != 0: + validation["result"] = "nonzero_exit" + detail = stderr.strip() or stdout.strip() or "no diagnostic output" + if len(detail) > 512: + detail = detail[:512] + "..." + raise RecordingError( + f"ros2 bag info failed with return code {returncode}: {detail}" + ) + validation["passed"] = True + validation["result"] = "passed" + return validation + + @staticmethod + def _output_text(value: object) -> str: + if value is None: + return "" + if isinstance(value, bytes): + return value.decode("utf-8", errors="replace") + return str(value) + + @staticmethod + def _write_bag_info_logs( + episode: _Episode, stdout: str, stderr: str + ) -> None: + # These diagnostics are intentionally not included in the data hashes; + # only metadata.yaml and MCAP payloads define a synchronized episode. + try: + (episode.active_directory / "bag_info.stdout.log").write_text( + stdout, encoding="utf-8" + ) + (episode.active_directory / "bag_info.stderr.log").write_text( + stderr, encoding="utf-8" + ) + except OSError: + pass + + def _preserve_failed_episode( + self, + episode: _Episode, + stop_reason: str, + error_text: str, + ) -> tuple[Path | None, str]: + if not episode.active_directory.exists(): + return None, "active episode directory is missing" + try: + ready_marker = episode.active_directory / "READY" + if ready_marker.exists(): + # Preserve, but invalidate, a marker if the final directory + # rename itself failed after the marker was written. + os.replace(ready_marker, episode.active_directory / "READY.invalid") + failed_manifest = { + "schema_version": 1, + "state": "failed", + "status": "failed", + "episode_id": episode.request.episode_id, + "session_id": episode.request.session_id, + "start_reason": episode.request.reason, + "stop_reason": stop_reason, + "topics": list(self.config.topics), + "required_topics": list(self.config.required_topics), + "required_topic_message_counts": ( + episode.required_topic_message_counts + ), + "storage_id": "mcap", + "custom_data": { + "capture_id": episode.request.episode_id, + "teleop_session_id": episode.request.session_id, + }, + "bag_info_validation": episode.bag_info_validation, + "started_at_utc": self._iso(episode.started_at), + "failed_at_utc": self._iso(self._as_utc(self._utc_now())), + "duration_seconds": max( + 0.0, self._monotonic() - episode.started_monotonic + ), + "error": error_text, + } + self._write_manifest(episode.active_directory, failed_manifest) + failed_directory = ( + Path(self.config.base_directory) + / "failed" + / episode.request.episode_id + ) + os.replace(episode.active_directory, failed_directory) + return failed_directory, "" + except Exception as archive_error: + return None, f"{type(archive_error).__name__}: {archive_error}" + + def _finish_status( + self, + result: str, + directory: Path | None, + error: str, + stop_reason: str, + ) -> None: + with self._condition: + self._status.update( + { + "state": "idle", + "recording": False, + "episode_id": None, + "session_id": None, + "started_at": None, + "stop_reason": stop_reason, + "last_result": result, + "last_episode_directory": ( + None if directory is None else str(directory) + ), + "last_error": error, + } + ) + self._requested_stop_reason = None + self._condition.notify_all() + + @staticmethod + def _close_logs(episode: _Episode) -> None: + for name in ("stdout_stream", "stderr_stream"): + stream = getattr(episode, name) + if stream is not None: + try: + stream.close() + except Exception: + # Recording data remains authoritative; a diagnostic-log + # close failure must not strand the episode in active/. + pass + finally: + setattr(episode, name, None) + + @staticmethod + def _file_record(root: Path, path: Path) -> dict[str, Any]: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + size = path.stat().st_size + return { + "path": path.relative_to(root).as_posix(), + # ``size`` is the on-wire sync protocol field. ``size_bytes`` is + # kept explicit for humans and status/manifest consumers. + "size": size, + "size_bytes": size, + "sha256": digest.hexdigest(), + } + + @staticmethod + def _write_manifest(directory: Path, manifest: Mapping[str, Any]) -> None: + temporary = directory / f".manifest.{uuid.uuid4().hex}.tmp" + final = directory / "manifest.json" + with temporary.open("w", encoding="utf-8") as stream: + json.dump(manifest, stream, ensure_ascii=False, indent=2, sort_keys=True) + stream.write("\n") + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary, final) + + @staticmethod + def _write_ready_marker(directory: Path) -> None: + """Publish READY last inside active before the atomic directory move.""" + + temporary = directory / f".READY.{uuid.uuid4().hex}.tmp" + final = directory / "READY" + with temporary.open("w", encoding="ascii") as stream: + stream.write("ready\n") + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary, final) + + @staticmethod + def _sync_directory(directory: Path) -> None: + descriptor = os.open(directory, os.O_RDONLY | getattr(os, "O_DIRECTORY", 0)) + try: + os.fsync(descriptor) + finally: + os.close(descriptor) + + @staticmethod + def _ros_node_name(episode_id: str) -> str: + digest = hashlib.sha256(episode_id.encode("utf-8")).hexdigest()[:16] + return f"tg3_data_recorder_{digest}" + + @staticmethod + def _validate_text(value: str, name: str) -> None: + if not isinstance(value, str) or not value.strip() or len(value) > 512: + raise ValueError( + f"{name} must be a non-empty string up to 512 characters" + ) + + @staticmethod + def _validate_custom_data_value(value: str, name: str) -> None: + allowed = ( + "-_.0123456789" + "ABCDEFGHIJKLMNOPQRSTUVWXYZ" + "abcdefghijklmnopqrstuvwxyz" + ) + if ( + not isinstance(value, str) + or not value + or len(value) > 128 + or any(character not in allowed for character in value) + ): + raise ValueError( + f"{name} must contain 1-128 safe custom-data characters" + ) + + @staticmethod + def _validate_episode_id(value: str) -> None: + allowed = ( + "-_.0123456789" + "ABCDEFGHIJKLMNOPQRSTUVWXYZ" + "abcdefghijklmnopqrstuvwxyz" + ) + first_allowed = ( + "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + ) + if ( + not isinstance(value, str) + or not value + or len(value) > 128 + or value[0] not in first_allowed + or value in (".", "..") + or Path(value).name != value + or any(character not in allowed for character in value) + ): + raise ValueError("episode ID must be one safe path component") + + @staticmethod + def _as_utc(value: datetime) -> datetime: + if not isinstance(value, datetime): + raise TypeError("UTC clock must return datetime") + if value.tzinfo is None: + raise ValueError("UTC clock must return a timezone-aware datetime") + return value.astimezone(timezone.utc) + + @staticmethod + def _iso(value: datetime) -> str: + return value.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") + + +__all__ = [ + "DataRecorderManager", + "RecorderConfig", + "RecordingError", + "RecordingToggleGate", + "ToggleAction", + "left_joystick_pressed", +] diff --git a/tg3_local_teleop/data_recorder_node.py b/tg3_local_teleop/data_recorder_node.py new file mode 100755 index 0000000..541a0f5 --- /dev/null +++ b/tg3_local_teleop/data_recorder_node.py @@ -0,0 +1,368 @@ +#!/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 + + +@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) + ) + + 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_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()) diff --git a/tg3_local_teleop/data_recorder_protocol.py b/tg3_local_teleop/data_recorder_protocol.py new file mode 100755 index 0000000..faa2e09 --- /dev/null +++ b/tg3_local_teleop/data_recorder_protocol.py @@ -0,0 +1,546 @@ +#!/usr/bin/env python3 +"""Pure control protocol for the independent TG3 data recorder. + +The ROS node is deliberately kept as a thin transport adapter. This module +owns JSON validation, idempotence, session/capture matching, acknowledgement +state, and the bridge-heartbeat watchdog, and can therefore be tested without +ROS installed. +""" + +from __future__ import annotations + +import json +import math +import time +from collections import OrderedDict +from collections.abc import Callable, Mapping, Sequence +from pathlib import Path +from typing import Any, Protocol + + +PROTOCOL_VERSION = 1 +ACTIVE_MANAGER_STATES = { + "start_pending", + "starting", + "recording", + "stop_pending", + "finalizing", +} + + +class RecorderBackend(Protocol): + """The non-blocking subset exposed by ``DataRecorderManager``.""" + + def request_start( + self, + session_id: str, + reason: str = "button", + *, + episode_id: str | None = None, + ) -> str | None: ... + + def request_stop(self, reason: str = "operator_button") -> bool: ... + + def status(self) -> dict[str, Any]: ... + + def shutdown( + self, *, wait: bool = True, timeout: float | None = None + ) -> bool: ... + + +class ProtocolError(ValueError): + """A control message is malformed or invalid for the current session.""" + + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + + +def _nonempty_text(value: Any, name: str, *, maximum: int = 128) -> str: + if ( + not isinstance(value, str) + or not value + or value.strip() != value + or len(value) > maximum + or any(ord(character) < 0x20 for character in value) + ): + raise ProtocolError( + f"invalid_{name}", + f"{name} must be non-empty text up to {maximum} characters", + ) + return value + + +def _capture_id(value: Any) -> str: + capture_id = _nonempty_text(value, "capture_id") + if ( + not capture_id[0].isalnum() + or Path(capture_id).name != capture_id + or capture_id in (".", "..") + or any( + character + not in "-_.0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + for character in capture_id + ) + ): + raise ProtocolError( + "invalid_capture_id", + "capture_id must be one safe path component beginning with a letter or digit", + ) + return capture_id + + +def _event_sequence(value: Any) -> int: + if type(value) is not int or value < 0 or value > (2**63 - 1): + raise ProtocolError( + "invalid_event_seq", + "event_seq must be an integer from 0 through 2^63-1", + ) + return value + + +def _optional_unix_time(value: Any) -> float | None: + if value is None: + return None + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ProtocolError( + "invalid_sent_unix_s", "sent_unix_s must be a finite number" + ) + result = float(value) + if not math.isfinite(result) or result < 0.0: + raise ProtocolError( + "invalid_sent_unix_s", "sent_unix_s must be a finite number" + ) + return result + + +def parse_control_message(text: str) -> dict[str, Any]: + """Parse one strict but forward-compatible recorder control object.""" + + if not isinstance(text, str) or not text or len(text.encode("utf-8")) > 16384: + raise ProtocolError( + "invalid_json", "control payload must be 1 through 16384 UTF-8 bytes" + ) + try: + raw = json.loads(text) + except (json.JSONDecodeError, UnicodeError) as error: + raise ProtocolError("invalid_json", f"invalid JSON: {error}") from error + if not isinstance(raw, Mapping): + raise ProtocolError("invalid_json_object", "control JSON must be an object") + if raw.get("version") != PROTOCOL_VERSION: + raise ProtocolError( + "unsupported_version", f"version must equal {PROTOCOL_VERSION}" + ) + + command = raw.get("command") + if command not in ("start", "stop", "heartbeat"): + raise ProtocolError( + "invalid_command", "command must be start, stop, or heartbeat" + ) + event_seq = _event_sequence(raw.get("event_seq")) + request_id = _nonempty_text(raw.get("request_id"), "request_id") + session_id = _nonempty_text( + raw.get("teleop_session_id"), "teleop_session_id" + ) + capture_id = _capture_id(raw.get("capture_id")) + sent_unix_s = _optional_unix_time(raw.get("sent_unix_s")) + + reason_value = raw.get("reason", "bridge_heartbeat" if command == "heartbeat" else "button") + reason = _nonempty_text(reason_value, "reason", maximum=512) + return { + "version": PROTOCOL_VERSION, + "command": command, + "event_seq": event_seq, + "request_id": request_id, + "teleop_session_id": session_id, + "capture_id": capture_id, + "reason": reason, + "sent_unix_s": sent_unix_s, + } + + +def topics_without_publishers( + required_topics: Sequence[str], + publisher_lookup: Callable[[str], Sequence[Any]], +) -> tuple[str, ...]: + """Return required topics that currently have no live publisher endpoint. + + ROS graph discovery can retain a topic name solely because this process is + subscribed to it. Looking up publisher endpoints is therefore the + meaningful start preflight; topic-name presence alone is insufficient. + Exceptions deliberately propagate so the caller can report graph-query + failure separately from an ordinary missing publisher. + """ + + missing: list[str] = [] + for topic in required_topics: + if not publisher_lookup(topic): + missing.append(topic) + return tuple(missing) + + +class RecorderControlProtocol: + """Idempotent bridge-to-recorder protocol and heartbeat watchdog.""" + + def __init__( + self, + recorder: RecorderBackend, + *, + heartbeat_timeout_seconds: float = 3.0, + monotonic: Callable[[], float] = time.monotonic, + unix_time: Callable[[], float] = time.time, + start_preflight: Callable[[], tuple[bool, str]] | None = None, + request_cache_size: int = 512, + session_cache_size: int = 64, + ) -> None: + if ( + not math.isfinite(heartbeat_timeout_seconds) + or heartbeat_timeout_seconds <= 0.0 + ): + raise ValueError("heartbeat timeout must be positive and finite") + if request_cache_size <= 0 or session_cache_size <= 0: + raise ValueError("protocol cache sizes must be positive") + self.recorder = recorder + self.heartbeat_timeout_seconds = float(heartbeat_timeout_seconds) + self._monotonic = monotonic + self._unix_time = unix_time + self._start_preflight = start_preflight + self._request_cache_size = request_cache_size + self._session_cache_size = session_cache_size + self._request_cache: OrderedDict[str, tuple[str, bool, str, int]] = ( + OrderedDict() + ) + self._last_sequence_by_session: OrderedDict[str, int] = OrderedDict() + self._capture_id: str | None = None + self._session_id: str | None = None + self._last_heartbeat_monotonic: float | None = None + self._watchdog_stop_requested = False + self._last_protocol_error = "" + self._ack_event_seq: int | None = None + self._ack_request_id: str | None = None + self._ack_accepted: bool | None = None + self._ack_code = "initialized" + self._accepted_count = 0 + self._rejected_count = 0 + self._duplicate_count = 0 + self._watchdog_stop_count = 0 + + def handle_json(self, text: str, *, now: float | None = None) -> dict[str, Any]: + """Handle one payload without blocking on recorder I/O.""" + + timestamp = self._monotonic() if now is None else float(now) + try: + event = parse_control_message(text) + except ProtocolError as error: + self._ack_event_seq = None + self._ack_request_id = None + self._ack_accepted = False + self._ack_code = error.code + self._last_protocol_error = str(error) + self._rejected_count += 1 + return self.status(now=timestamp) + + fingerprint = json.dumps(event, sort_keys=True, separators=(",", ":")) + request_id = event["request_id"] + cached = self._request_cache.get(request_id) + if cached is not None: + cached_fingerprint, accepted, code, event_seq = cached + if cached_fingerprint != fingerprint: + return self._reject( + event, + "request_id_reused", + "request_id was already used for a different payload", + fingerprint=fingerprint, + cache=False, + now=timestamp, + ) + self._request_cache.move_to_end(request_id) + self._duplicate_count += 1 + self._set_ack(event_seq, request_id, accepted, code) + return self.status(now=timestamp) + + session_id = event["teleop_session_id"] + last_sequence = self._last_sequence_by_session.get(session_id) + if last_sequence is not None and event["event_seq"] <= last_sequence: + return self._reject( + event, + "stale_event_seq", + f"event_seq must be greater than the previous value {last_sequence}", + fingerprint=fingerprint, + now=timestamp, + ) + + self._remember_sequence(session_id, event["event_seq"]) + self._last_protocol_error = "" + try: + accepted, code = self._execute(event, timestamp) + message = ( + "" + if accepted + else (self._last_protocol_error or code) + ) + except Exception as error: + accepted = False + code = "recorder_exception" + message = f"{type(error).__name__}: {error}" + + if accepted: + self._accepted_count += 1 + self._last_protocol_error = "" + else: + self._rejected_count += 1 + self._last_protocol_error = message + self._remember_request( + request_id, + fingerprint, + accepted, + code, + event["event_seq"], + ) + self._set_ack( + event["event_seq"], + request_id, + accepted, + code, + ) + return self.status(now=timestamp) + + def poll(self, *, now: float | None = None) -> dict[str, Any]: + """Advance the watchdog and return the current status snapshot.""" + + timestamp = self._monotonic() if now is None else float(now) + recorder_status = self._safe_recorder_status() + manager_state = str(recorder_status.get("state", "unknown")) + if ( + manager_state in ("start_pending", "starting", "recording") + and self._last_heartbeat_monotonic is not None + and not self._watchdog_stop_requested + and timestamp - self._last_heartbeat_monotonic + > self.heartbeat_timeout_seconds + ): + try: + accepted = self.recorder.request_stop("bridge_heartbeat_timeout") + except Exception as error: + self._last_protocol_error = ( + f"watchdog stop failed: {type(error).__name__}: {error}" + ) + else: + if accepted: + self._watchdog_stop_requested = True + self._watchdog_stop_count += 1 + return self._status_from_snapshot(recorder_status, timestamp) + + def status(self, *, now: float | None = None) -> dict[str, Any]: + timestamp = self._monotonic() if now is None else float(now) + return self._status_from_snapshot(self._safe_recorder_status(), timestamp) + + def shutdown(self, *, timeout: float | None = None) -> bool: + """Request clean recorder finalization when the supervisor exits.""" + + return self.recorder.shutdown(wait=True, timeout=timeout) + + def _execute(self, event: Mapping[str, Any], now: float) -> tuple[bool, str]: + command = event["command"] + if command == "start": + return self._start(event, now) + if command == "stop": + return self._stop(event) + return self._heartbeat(event, now) + + def _start(self, event: Mapping[str, Any], now: float) -> tuple[bool, str]: + recorder_status = self._safe_recorder_status() + manager_state = str(recorder_status.get("state", "unknown")) + if manager_state in ACTIVE_MANAGER_STATES: + if self._matches_context(event): + return True, "already_active" + return False, "recorder_busy" + if manager_state not in ("idle",): + return False, "recorder_unavailable" + if self._start_preflight is not None: + allowed, detail = self._start_preflight() + if not allowed: + self._last_protocol_error = detail + return False, "preflight_failed" + + actual_capture_id = self.recorder.request_start( + event["teleop_session_id"], + event["reason"], + episode_id=event["capture_id"], + ) + if actual_capture_id is None: + return False, "recorder_busy" + if actual_capture_id != event["capture_id"]: + # A mismatched directory would make bridge samples and manifests + # impossible to correlate, so immediately fail safe by stopping. + self.recorder.request_stop("capture_id_mismatch") + return False, "capture_id_mismatch" + self._capture_id = actual_capture_id + self._session_id = event["teleop_session_id"] + self._last_heartbeat_monotonic = now + self._watchdog_stop_requested = False + return True, "start_accepted" + + def _stop(self, event: Mapping[str, Any]) -> tuple[bool, str]: + recorder_status = self._safe_recorder_status() + manager_state = str(recorder_status.get("state", "unknown")) + if manager_state not in ACTIVE_MANAGER_STATES: + if self._matches_context(event) or self._capture_id is None: + return True, "already_stopped" + return False, "capture_not_active" + if not self._matches_context(event): + return False, "capture_mismatch" + accepted = self.recorder.request_stop(event["reason"]) + if not accepted: + return False, "stop_rejected" + self._watchdog_stop_requested = True + return True, "stop_accepted" + + def _heartbeat( + self, event: Mapping[str, Any], now: float + ) -> tuple[bool, str]: + if not self._matches_context(event): + return False, "capture_mismatch" + recorder_status = self._safe_recorder_status() + manager_state = str(recorder_status.get("state", "unknown")) + if manager_state not in ACTIVE_MANAGER_STATES: + return True, "already_stopped" + self._last_heartbeat_monotonic = now + return True, "heartbeat_accepted" + + def _matches_context(self, event: Mapping[str, Any]) -> bool: + return ( + self._capture_id == event["capture_id"] + and self._session_id == event["teleop_session_id"] + ) + + def _reject( + self, + event: Mapping[str, Any], + code: str, + message: str, + *, + fingerprint: str, + cache: bool = True, + now: float, + ) -> dict[str, Any]: + self._rejected_count += 1 + self._last_protocol_error = message + self._set_ack( + event["event_seq"], event["request_id"], False, code + ) + if cache: + self._remember_request( + event["request_id"], + fingerprint, + False, + code, + event["event_seq"], + ) + return self.status(now=now) + + def _remember_sequence(self, session_id: str, sequence: int) -> None: + self._last_sequence_by_session[session_id] = sequence + self._last_sequence_by_session.move_to_end(session_id) + while len(self._last_sequence_by_session) > self._session_cache_size: + self._last_sequence_by_session.popitem(last=False) + + def _remember_request( + self, + request_id: str, + fingerprint: str, + accepted: bool, + code: str, + event_seq: int, + ) -> None: + self._request_cache[request_id] = ( + fingerprint, + accepted, + code, + event_seq, + ) + self._request_cache.move_to_end(request_id) + while len(self._request_cache) > self._request_cache_size: + self._request_cache.popitem(last=False) + + def _set_ack( + self, + event_seq: int, + request_id: str, + accepted: bool, + code: str, + ) -> None: + self._ack_event_seq = event_seq + self._ack_request_id = request_id + self._ack_accepted = accepted + self._ack_code = code + + def _safe_recorder_status(self) -> dict[str, Any]: + try: + status = self.recorder.status() + except Exception as error: + self._last_protocol_error = ( + f"recorder status failed: {type(error).__name__}: {error}" + ) + return { + "state": "unknown", + "recording": False, + "last_error": self._last_protocol_error, + } + return dict(status) + + def _status_from_snapshot( + self, recorder_status: Mapping[str, Any], now: float + ) -> dict[str, Any]: + manager_state = str(recorder_status.get("state", "unknown")) + if manager_state in ("start_pending", "starting"): + state = "starting" + elif manager_state == "recording": + state = "recording" + elif manager_state in ("stop_pending", "finalizing"): + state = "stopping" + elif manager_state == "idle": + last_result = recorder_status.get("last_result") + state = last_result if last_result in ("ready", "failed") else "idle" + else: + state = "failed" + + heartbeat_age: float | None = None + if self._last_heartbeat_monotonic is not None: + heartbeat_age = max(0.0, now - self._last_heartbeat_monotonic) + manager_error = recorder_status.get("last_error") + last_error = ( + str(manager_error) + if isinstance(manager_error, str) and manager_error + else self._last_protocol_error + ) + return { + "version": PROTOCOL_VERSION, + "state": state, + "capture_id": self._capture_id, + "teleop_session_id": self._session_id, + "ack_event_seq": self._ack_event_seq, + "ack_request_id": self._ack_request_id, + "ack_accepted": self._ack_accepted, + "ack_code": self._ack_code, + "last_error": last_error, + "stop_reason": recorder_status.get("stop_reason"), + "heartbeat_age_s": heartbeat_age, + "heartbeat_timeout_s": self.heartbeat_timeout_seconds, + "updated_unix_s": float(self._unix_time()), + "statistics": { + "accepted": self._accepted_count, + "rejected": self._rejected_count, + "duplicates": self._duplicate_count, + "watchdog_stops": self._watchdog_stop_count, + }, + "recorder": dict(recorder_status), + } + + +__all__ = [ + "PROTOCOL_VERSION", + "ProtocolError", + "RecorderControlProtocol", + "parse_control_message", + "topics_without_publishers", +] diff --git a/tg3_local_teleop/run_data_recorder.sh b/tg3_local_teleop/run_data_recorder.sh new file mode 100755 index 0000000..47be1bb --- /dev/null +++ b/tg3_local_teleop/run_data_recorder.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash +# Load every overlay needed to deserialize the TianGong custom message types +# before rosbag2 discovers the configured topics. +set -eo pipefail + +APP_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +required_setups=( + /opt/ros/jazzy/setup.bash + /home/nvidia/xos/setup.bash + /opt/robot_tele_server/install/setup.bash + "$APP_DIR/ros2_py/install/setup.bash" +) +for setup_file in "${required_setups[@]}"; do + if [[ ! -f "$setup_file" ]]; then + echo "Required ROS setup is missing: $setup_file" >&2 + exit 1 + fi + # shellcheck disable=SC1090 + source "$setup_file" +done +set -u + +exec python3 "$APP_DIR/data_recorder_node.py" \ + --config "$APP_DIR/config.toml" \ + "$@" diff --git a/tg3_local_teleop/test_data_collection.py b/tg3_local_teleop/test_data_collection.py new file mode 100755 index 0000000..330e090 --- /dev/null +++ b/tg3_local_teleop/test_data_collection.py @@ -0,0 +1,699 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import hashlib +import json +import signal +import subprocess +import tempfile +import threading +import time +import unittest +from collections.abc import Callable, Mapping, Sequence +from pathlib import Path +from typing import Any + +from data_collection import ( + DataRecorderManager, + RecorderConfig, + RecordingToggleGate, + left_joystick_pressed, +) + + +class LeftJoystickParserTest(unittest.TestCase): + def test_accepts_only_live_bool_or_binary_integer(self) -> None: + self.assertIs( + left_joystick_pressed({"button_joystick": {"left": True}}), + True, + ) + self.assertIs( + left_joystick_pressed({"button_joystick": {"left": False}}), + False, + ) + self.assertIs( + left_joystick_pressed({"button_joystick": {"left": 1}}), True + ) + self.assertIs( + left_joystick_pressed({"button_joystick": {"left": 0}}), False + ) + for value in (-1, 2, None, "true", [], 1.0, {"pressed": True}): + with self.subTest(value=value): + self.assertIsNone( + left_joystick_pressed({"button_joystick": {"left": value}}) + ) + + def test_rejects_every_malformed_container(self) -> None: + malformed: list[Any] = [ + {}, + {"button_joystick": None}, + {"button_joystick": []}, + {"button_joystick": {}}, + ] + for sample in malformed: + with self.subTest(sample=sample): + self.assertIsNone(left_joystick_pressed(sample)) + + +class RecordingToggleGateTest(unittest.TestCase): + def setUp(self) -> None: + self.gate = RecordingToggleGate( + hold_seconds=1.0, release_seconds=0.5 + ) + self.gate.new_session() + + def stable_release(self, now: float) -> float: + self.assertIsNone( + self.gate.update(now, input_healthy=True, pressed=False) + ) + now += 0.5 + self.assertIsNone( + self.gate.update(now, input_healthy=True, pressed=False) + ) + self.assertFalse(self.gate.require_release) + return now + + def hold(self, now: float) -> tuple[float, str | None]: + self.assertIsNone( + self.gate.update(now, input_healthy=True, pressed=True) + ) + now += 1.0 + return now, self.gate.update(now, input_healthy=True, pressed=True) + + def test_one_button_starts_and_stops_once_per_physical_hold(self) -> None: + now = self.stable_release(0.0) + now, action = self.hold(now + 0.01) + self.assertEqual(action, "start") + self.assertTrue(self.gate.active) + + # Continuing the same press for any duration cannot stop recording. + self.assertIsNone( + self.gate.update(now + 20.0, input_healthy=True, pressed=True) + ) + self.assertTrue(self.gate.active) + + now = self.stable_release(now + 20.01) + now, action = self.hold(now + 0.01) + self.assertEqual(action, "stop") + self.assertFalse(self.gate.active) + self.assertEqual(self.gate.toggle_count, 2) + + def test_new_session_is_release_locked(self) -> None: + self.assertIsNone( + self.gate.update(0.0, input_healthy=True, pressed=True) + ) + self.assertIsNone( + self.gate.update(100.0, input_healthy=True, pressed=True) + ) + self.assertFalse(self.gate.active) + self.assertEqual(self.gate.state, "awaiting_release") + + def test_invalid_input_never_counts_as_release(self) -> None: + self.gate.update(0.0, input_healthy=True, pressed=False) + self.gate.update(0.49, input_healthy=True, pressed=False) + self.gate.update(0.5, input_healthy=True, pressed=None) + self.assertTrue(self.gate.require_release) + self.gate.update(10.0, input_healthy=True, pressed=False) + self.gate.update(10.49, input_healthy=True, pressed=False) + self.assertTrue(self.gate.require_release) + self.gate.update(10.5, input_healthy=True, pressed=False) + self.assertFalse(self.gate.require_release) + + self.gate.update(10.6, input_healthy=True, pressed=True) + self.gate.update(11.59, input_healthy=False, pressed=True) + self.gate.update(20.0, input_healthy=True, pressed=True) + self.gate.update(30.0, input_healthy=True, pressed=True) + self.assertFalse(self.gate.active) + self.assertTrue(self.gate.require_release) + + def test_short_press_requires_another_stable_release(self) -> None: + now = self.stable_release(0.0) + self.gate.update(now + 0.1, input_healthy=True, pressed=True) + self.gate.update(now + 0.9, input_healthy=True, pressed=False) + self.assertTrue(self.gate.require_release) + self.gate.update(now + 1.39, input_healthy=True, pressed=False) + self.assertTrue(self.gate.require_release) + self.gate.update(now + 1.4, input_healthy=True, pressed=False) + self.assertFalse(self.gate.require_release) + + def test_session_end_reports_required_stop_and_relocks(self) -> None: + now = self.stable_release(0.0) + _, action = self.hold(now + 0.1) + self.assertEqual(action, "start") + self.assertEqual(self.gate.end_session(), "stop") + self.assertEqual(self.gate.state, "closed") + self.gate.new_session() + self.assertTrue(self.gate.require_release) + self.assertFalse(self.gate.active) + + def test_force_inactive_handles_automatic_recorder_stop(self) -> None: + now = self.stable_release(0.0) + self.hold(now + 0.1) + self.gate.force_inactive("maximum_duration") + self.assertFalse(self.gate.active) + self.assertTrue(self.gate.require_release) + self.assertEqual(self.gate.last_transition, "maximum_duration") + + +class FakeProcess: + def __init__( + self, + command: Sequence[str], + *, + valid_bag: bool = True, + empty_mcap: bool = False, + unexpected_returncode: int | None = None, + ignore_sigint: bool = False, + topic_counts: Mapping[str, int] | None = None, + ) -> None: + self.command = list(command) + self.signals: list[int] = [] + self.returncode = unexpected_returncode + self.ignore_sigint = ignore_sigint + output = Path(self.command[self.command.index("--output") + 1]) + output.mkdir(parents=True) + if valid_bag: + if topic_counts is None: + topic_counts = {"/joint_states": 25, "/tf": 50} + metadata = { + "rosbag2_bagfile_information": { + "storage_identifier": "mcap", + "topics_with_message_count": [ + { + "topic_metadata": {"name": topic}, + "message_count": count, + } + for topic, count in topic_counts.items() + ], + } + } + (output / "metadata.yaml").write_text( + json.dumps(metadata), + encoding="utf-8", + ) + (output / "data_0.mcap").write_bytes( + b"" if empty_mcap else b"fake-mcap-payload" + ) + + def poll(self) -> int | None: + return self.returncode + + def send_signal(self, sig: int) -> None: + self.signals.append(sig) + if not self.ignore_sigint: + self.returncode = 0 + + def wait(self, timeout: float | None = None) -> int: + if self.returncode is None: + raise subprocess.TimeoutExpired(self.command, timeout) + return self.returncode + + def terminate(self) -> None: + self.signals.append(signal.SIGTERM) + + def kill(self) -> None: + self.signals.append(signal.SIGKILL) + self.returncode = -signal.SIGKILL + + +class FakeProcessFactory: + def __init__(self, **process_options: Any) -> None: + self.process_options = process_options + self.processes: list[FakeProcess] = [] + self.commands: list[list[str]] = [] + self.kwargs: list[dict[str, Any]] = [] + + def __call__(self, command: Sequence[str], **kwargs: Any) -> FakeProcess: + self.commands.append(list(command)) + self.kwargs.append(kwargs) + process = FakeProcess(command, **self.process_options) + self.processes.append(process) + return process + + +class FakeBagInfoRunner: + def __init__( + self, + *, + returncode: int = 0, + stdout: str = "Files: data_0.mcap\n", + stderr: str = "", + timeout: bool = False, + ) -> None: + self.returncode = returncode + self.stdout = stdout + self.stderr = stderr + self.timeout = timeout + self.calls: list[tuple[list[str], dict[str, Any]]] = [] + + def __call__(self, command: Sequence[str], **kwargs: Any) -> Any: + command_list = list(command) + self.calls.append((command_list, kwargs)) + if self.timeout: + raise subprocess.TimeoutExpired( + command_list, + kwargs.get("timeout"), + output="partial bag info", + stderr="timed out", + ) + return subprocess.CompletedProcess( + command_list, + self.returncode, + stdout=self.stdout, + stderr=self.stderr, + ) + + +class DataRecorderManagerTest(unittest.TestCase): + def setUp(self) -> None: + self.temporary = tempfile.TemporaryDirectory() + self.base = Path(self.temporary.name) / "Data_Get" + self.managers: list[DataRecorderManager] = [] + self.episode_number = 0 + + def tearDown(self) -> None: + for manager in self.managers: + manager.shutdown(wait=True, timeout=2.0) + self.temporary.cleanup() + + def make_manager( + self, + factory: Callable[..., FakeProcess] | None = None, + *, + free_bytes: Callable[[Path], int] | None = None, + max_duration: float = 30.0, + minimum_free: int = 100, + process_options: dict[str, Any] | None = None, + command_runner: Callable[..., Any] | None = None, + validate_bag_info: bool = True, + required_topics: Sequence[str] = ("/joint_states", "/tf"), + ) -> tuple[DataRecorderManager, Any]: + if factory is None: + factory = FakeProcessFactory(**(process_options or {})) + + def episode_id() -> str: + self.episode_number += 1 + return f"episode-{self.episode_number:03d}" + + manager = DataRecorderManager( + RecorderConfig( + base_directory=self.base, + topics=("/joint_states", "/tf"), + required_topics=required_topics, + minimum_free_bytes=minimum_free, + max_duration_seconds=max_duration, + poll_interval_seconds=0.005, + sigint_timeout_seconds=0.01, + kill_timeout_seconds=0.01, + validate_bag_info=validate_bag_info, + bag_info_timeout_seconds=0.02, + ), + process_factory=factory, + command_runner=command_runner or FakeBagInfoRunner(), + free_bytes=free_bytes or (lambda _path: 10_000), + episode_id_factory=episode_id, + ) + self.managers.append(manager) + return manager, factory + + def wait_for( + self, predicate: Callable[[], bool], timeout: float = 2.0 + ) -> None: + deadline = time.monotonic() + timeout + while not predicate(): + if time.monotonic() >= deadline: + self.fail("condition did not become true before timeout") + time.sleep(0.002) + + def test_manual_stop_uses_sigint_and_atomically_creates_ready_manifest( + self, + ) -> None: + bag_info = FakeBagInfoRunner() + manager, factory = self.make_manager(command_runner=bag_info) + self.assertTrue(manager.request_start("a" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + status = manager.status() + self.assertEqual(status["state"], "recording") + self.assertTrue(manager.request_stop("operator_button")) + self.assertTrue(manager.wait_until_idle(2.0)) + + status = manager.status() + self.assertEqual(status["last_result"], "ready") + self.assertEqual(status["stop_reason"], "operator_button") + ready = Path(status["last_episode_directory"]) + self.assertEqual(ready.parent, self.base / "ready") + self.assertFalse((self.base / "active" / ready.name).exists()) + manifest = json.loads((ready / "manifest.json").read_text()) + self.assertEqual(manifest["session_id"], "a" * 32) + self.assertEqual(manifest["topics"], ["/joint_states", "/tf"]) + self.assertEqual( + manifest["required_topic_message_counts"], + {"/joint_states": 25, "/tf": 50}, + ) + self.assertEqual(manifest["storage_id"], "mcap") + self.assertEqual(manifest["state"], "complete") + self.assertEqual(manifest["stop_reason"], "operator_button") + self.assertTrue((ready / "READY").is_file()) + self.assertEqual((ready / "READY").read_text(), "ready\n") + self.assertEqual( + manifest["bag_info_validation"]["result"], "passed" + ) + self.assertTrue(manifest["bag_info_validation"]["passed"]) + self.assertEqual( + manifest["custom_data"], + { + "capture_id": ready.name, + "teleop_session_id": "a" * 32, + }, + ) + self.assertIn("started_at_utc", manifest) + self.assertIn("stopped_at_utc", manifest) + paths = {record["path"]: record for record in manifest["files"]} + self.assertEqual( + set(paths), {"bag/metadata.yaml", "bag/data_0.mcap"} + ) + mcap_record = paths["bag/data_0.mcap"] + payload = (ready / "bag/data_0.mcap").read_bytes() + self.assertEqual(mcap_record["size_bytes"], len(payload)) + self.assertEqual(mcap_record["size"], len(payload)) + self.assertEqual(mcap_record["sha256"], hashlib.sha256(payload).hexdigest()) + self.assertEqual(factory.processes[0].signals, [signal.SIGINT]) + command = factory.commands[0] + self.assertEqual( + command[:5], ["ros2", "bag", "record", "--storage", "mcap"] + ) + self.assertEqual( + command[command.index("--storage-preset-profile") + 1], + "zstd_fast", + ) + self.assertEqual(command[command.index("--max-cache-size") + 1], "67108864") + self.assertEqual(command[command.index("--max-bag-duration") + 1], "300") + self.assertIn("--disable-keyboard-controls", command) + self.assertRegex( + command[command.index("--node-name") + 1], + r"^tg3_data_recorder_[0-9a-f]{16}$", + ) + self.assertIn("--topics", command) + custom_index = command.index("--custom-data") + topics_index = command.index("--topics") + self.assertLess(custom_index, topics_index) + self.assertEqual( + command[custom_index + 1 : custom_index + 3], + [f"capture_id={ready.name}", f"teleop_session_id={'a' * 32}"], + ) + self.assertEqual(command[-2:], ["/joint_states", "/tf"]) + self.assertTrue(factory.kwargs[0]["start_new_session"]) + self.assertEqual( + bag_info.calls[0][0], + [ + "ros2", + "bag", + "info", + str(ready.parent.parent / "active" / ready.name / "bag"), + ], + ) + self.assertEqual(bag_info.calls[0][1]["timeout"], 0.02) + self.assertTrue(bag_info.calls[0][1]["text"]) + self.assertFalse(bag_info.calls[0][1]["check"]) + + def test_request_methods_do_not_wait_for_slow_process_factory(self) -> None: + entered = threading.Event() + release = threading.Event() + wrapped_factory = FakeProcessFactory() + + def slow_factory(command: Sequence[str], **kwargs: Any) -> FakeProcess: + entered.set() + release.wait(1.0) + return wrapped_factory(command, **kwargs) + + manager, _ = self.make_manager(factory=slow_factory) + started = time.monotonic() + self.assertTrue(manager.request_start("b" * 32)) + self.assertLess(time.monotonic() - started, 0.05) + self.assertTrue(entered.wait(1.0)) + started = time.monotonic() + self.assertTrue(manager.request_stop("operator_button")) + self.assertLess(time.monotonic() - started, 0.05) + release.set() + self.assertTrue(manager.wait_until_idle(2.0)) + + def test_low_disk_space_automatically_stops_valid_episode(self) -> None: + calls = 0 + + def disk(_path: Path) -> int: + nonlocal calls + calls += 1 + return 10_000 if calls == 1 else 0 + + manager, factory = self.make_manager(free_bytes=disk) + self.assertTrue(manager.request_start("c" * 32)) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "ready") + self.assertEqual(status["stop_reason"], "low_disk_space") + self.assertIn(signal.SIGINT, factory.processes[0].signals) + + def test_maximum_duration_automatically_stops(self) -> None: + manager, _ = self.make_manager(max_duration=0.02) + self.assertTrue(manager.request_start("d" * 32)) + self.assertTrue(manager.wait_until_idle(2.0)) + self.assertEqual(manager.status()["last_result"], "ready") + self.assertEqual(manager.status()["stop_reason"], "maximum_duration") + + def test_missing_or_empty_mcap_is_preserved_in_failed(self) -> None: + for options in ({"valid_bag": False}, {"empty_mcap": True}): + with self.subTest(options=options): + manager, _ = self.make_manager(process_options=options) + self.assertTrue(manager.request_start("e" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + failed = Path(status["last_episode_directory"]) + self.assertEqual(failed.parent, self.base / "failed") + failure = json.loads((failed / "manifest.json").read_text()) + self.assertEqual(failure["status"], "failed") + self.assertEqual(failure["session_id"], "e" * 32) + self.assertTrue(failure["error"]) + manager.shutdown(timeout=2.0) + + def test_unexpected_ros_exit_is_failed_not_ready(self) -> None: + manager, _ = self.make_manager( + process_options={"unexpected_returncode": 7} + ) + self.assertTrue(manager.request_start("f" * 32)) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("return code 7", status["last_error"]) + self.assertEqual( + Path(status["last_episode_directory"]).parent, + self.base / "failed", + ) + + def test_bag_info_nonzero_exit_preserves_episode_as_failed(self) -> None: + bag_info = FakeBagInfoRunner(returncode=4, stderr="MCAP read failed") + manager, _ = self.make_manager(command_runner=bag_info) + self.assertTrue(manager.request_start("7" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("return code 4", status["last_error"]) + failed = Path(status["last_episode_directory"]) + manifest = json.loads((failed / "manifest.json").read_text()) + validation = manifest["bag_info_validation"] + self.assertFalse(validation["passed"]) + self.assertEqual(validation["result"], "nonzero_exit") + self.assertEqual(validation["returncode"], 4) + self.assertFalse((failed / "READY").exists()) + + def test_missing_required_topic_is_failed_before_bag_info(self) -> None: + bag_info = FakeBagInfoRunner() + manager, _ = self.make_manager( + command_runner=bag_info, + process_options={"topic_counts": {"/joint_states": 25}}, + ) + self.assertTrue(manager.request_start("0" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("missing required topics: /tf", status["last_error"]) + failed = Path(status["last_episode_directory"]) + manifest = json.loads((failed / "manifest.json").read_text()) + self.assertEqual( + manifest["required_topic_message_counts"], + {"/joint_states": 25, "/tf": 0}, + ) + self.assertEqual(bag_info.calls, []) + self.assertFalse((failed / "READY").exists()) + + def test_zero_message_required_topic_is_failed_before_bag_info(self) -> None: + bag_info = FakeBagInfoRunner() + manager, _ = self.make_manager( + command_runner=bag_info, + process_options={ + "topic_counts": {"/joint_states": 25, "/tf": 0} + }, + ) + self.assertTrue(manager.request_start("a0" * 16)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("zero-message required topics: /tf", status["last_error"]) + failed = Path(status["last_episode_directory"]) + manifest = json.loads((failed / "manifest.json").read_text()) + self.assertEqual(manifest["required_topic_message_counts"]["/tf"], 0) + self.assertEqual(bag_info.calls, []) + self.assertFalse((failed / "READY").exists()) + + def test_bag_info_timeout_preserves_episode_as_failed(self) -> None: + manager, _ = self.make_manager( + command_runner=FakeBagInfoRunner(timeout=True) + ) + self.assertTrue(manager.request_start("8" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("timed out", status["last_error"]) + manifest = json.loads( + ( + Path(status["last_episode_directory"]) / "manifest.json" + ).read_text() + ) + self.assertEqual( + manifest["bag_info_validation"]["result"], "timeout" + ) + + def test_bag_info_validation_can_be_explicitly_disabled(self) -> None: + bag_info = FakeBagInfoRunner() + manager, _ = self.make_manager( + command_runner=bag_info, validate_bag_info=False + ) + self.assertTrue(manager.request_start("9" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + ready = Path(manager.status()["last_episode_directory"]) + manifest = json.loads((ready / "manifest.json").read_text()) + self.assertEqual( + manifest["bag_info_validation"]["result"], "disabled" + ) + self.assertEqual(bag_info.calls, []) + + def test_sigint_timeout_forces_kill_and_marks_failed(self) -> None: + manager, factory = self.make_manager( + process_options={"ignore_sigint": True} + ) + self.assertTrue(manager.request_start("1" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + self.assertEqual(manager.status()["last_result"], "failed") + signals = factory.processes[0].signals + self.assertEqual(signals[0], signal.SIGINT) + self.assertIn(signal.SIGTERM, signals) + self.assertIn(signal.SIGKILL, signals) + + def test_shutdown_waits_for_active_recording_to_finalize(self) -> None: + manager, factory = self.make_manager() + self.assertTrue(manager.request_start("2" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.shutdown(wait=True, timeout=2.0)) + status = manager.status() + self.assertEqual(status["state"], "shutdown") + self.assertEqual(status["last_result"], "ready") + self.assertEqual(status["stop_reason"], "shutdown") + self.assertIn(signal.SIGINT, factory.processes[0].signals) + + def test_status_snapshot_is_a_copy_and_start_is_single_flight(self) -> None: + manager, _ = self.make_manager() + first = manager.status() + first["topics"].append("/mutated") + first["state"] = "corrupt" + self.assertNotIn("/mutated", manager.status()["topics"]) + self.assertEqual(manager.status()["state"], "idle") + self.assertTrue(manager.request_start("3" * 32)) + self.assertFalse(manager.request_start("4" * 32)) + self.wait_for(lambda: manager.status()["recording"]) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + + def test_caller_capture_id_is_validated_and_returned(self) -> None: + manager, _ = self.make_manager() + capture_id = "capture-20260810T120000Z-abc123" + self.assertEqual( + manager.request_start("6" * 32, episode_id=capture_id), + capture_id, + ) + self.wait_for(lambda: manager.status()["recording"]) + self.assertEqual(manager.status()["episode_id"], capture_id) + self.assertTrue(manager.request_stop()) + self.assertTrue(manager.wait_until_idle(2.0)) + self.assertEqual( + Path(manager.status()["last_episode_directory"]).name, + capture_id, + ) + with self.assertRaises(ValueError): + manager.request_start("6" * 32, episode_id="../escape") + with self.assertRaises(ValueError): + manager.request_start("6" * 32, episode_id=".hidden") + with self.assertRaises(ValueError): + manager.request_start("6" * 32, episode_id="x" * 129) + with self.assertRaises(ValueError): + manager.request_start("unsafe session", episode_id="capture-safe") + + def test_low_disk_before_spawn_reports_failed_without_starting_ros(self) -> None: + manager, factory = self.make_manager(free_bytes=lambda _path: 0) + self.assertTrue(manager.request_start("5" * 32)) + self.assertTrue(manager.wait_until_idle(2.0)) + status = manager.status() + self.assertEqual(status["last_result"], "failed") + self.assertIn("insufficient free space", status["last_error"]) + self.assertEqual(factory.processes, []) + self.assertEqual( + Path(status["last_episode_directory"]).parent, + self.base / "failed", + ) + failure = json.loads( + (Path(status["last_episode_directory"]) / "manifest.json").read_text() + ) + self.assertEqual(failure["status"], "failed") + + def test_configuration_rejects_unsafe_or_ambiguous_values(self) -> None: + with self.assertRaises(ValueError): + RecorderConfig(self.base, ()) + with self.assertRaises(ValueError): + RecorderConfig(self.base, ("relative",)) + with self.assertRaises(ValueError): + RecorderConfig(self.base, ("/same", "/same")) + with self.assertRaises(ValueError): + RecorderConfig(self.base, ("/ok",), minimum_free_bytes=-1) + with self.assertRaises(ValueError): + RecorderConfig(self.base, ("/ok",), validate_bag_info=1) + with self.assertRaises(ValueError): + RecorderConfig(self.base, ("/ok",), bag_info_timeout_seconds=0) + with self.assertRaises(ValueError): + RecorderConfig( + self.base, + ("/recorded",), + required_topics=("/not_recorded",), + ) + with self.assertRaises(ValueError): + RecorderConfig( + self.base, + ("/recorded",), + required_topics=("/recorded", "/recorded"), + ) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tg3_local_teleop/test_data_recorder_protocol.py b/tg3_local_teleop/test_data_recorder_protocol.py new file mode 100755 index 0000000..d83f684 --- /dev/null +++ b/tg3_local_teleop/test_data_recorder_protocol.py @@ -0,0 +1,274 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import json +import unittest +from typing import Any + +from data_recorder_protocol import ( + ProtocolError, + RecorderControlProtocol, + parse_control_message, + topics_without_publishers, +) + + +class FakeRecorder: + def __init__(self) -> None: + self.snapshot: dict[str, Any] = { + "state": "idle", + "recording": False, + "episode_id": None, + "session_id": None, + "stop_reason": None, + "last_result": None, + "last_error": "", + "topics": ["/robot_state"], + } + self.starts: list[tuple[str, str, str | None]] = [] + self.stops: list[str] = [] + self.shutdowns = 0 + self.accept_start = True + + def request_start( + self, + session_id: str, + reason: str = "button", + *, + episode_id: str | None = None, + ) -> str | None: + self.starts.append((session_id, reason, episode_id)) + if not self.accept_start or self.snapshot["state"] != "idle": + return None + self.snapshot.update( + { + "state": "start_pending", + "episode_id": episode_id, + "session_id": session_id, + "last_result": None, + } + ) + return episode_id + + def request_stop(self, reason: str = "operator_button") -> bool: + self.stops.append(reason) + if self.snapshot["state"] not in ( + "start_pending", + "starting", + "recording", + "stop_pending", + ): + return False + self.snapshot["state"] = "stop_pending" + self.snapshot["stop_reason"] = reason + return True + + def status(self) -> dict[str, Any]: + return dict(self.snapshot) + + def shutdown( + self, *, wait: bool = True, timeout: float | None = None + ) -> bool: + self.shutdowns += 1 + return True + + +def control( + command: str, + sequence: int, + request_id: str, + *, + capture_id: str = "capture-001", + session_id: str = "teleop-session-001", + reason: str = "test", +) -> str: + return json.dumps( + { + "version": 1, + "command": command, + "event_seq": sequence, + "request_id": request_id, + "capture_id": capture_id, + "teleop_session_id": session_id, + "reason": reason, + "sent_unix_s": 1000.0 + sequence, + } + ) + + +class ParseControlMessageTest(unittest.TestCase): + def test_accepts_forward_compatible_object(self) -> None: + raw = json.loads(control("start", 0, "request-0")) + raw["future_field"] = {"ignored": True} + parsed = parse_control_message(json.dumps(raw)) + self.assertEqual(parsed["command"], "start") + self.assertEqual(parsed["capture_id"], "capture-001") + + def test_rejects_unsafe_or_ambiguous_values(self) -> None: + bad_objects = [ + [], + {"version": 2}, + { + "version": 1, + "command": "erase", + "event_seq": 1, + "request_id": "r", + "capture_id": "c", + "teleop_session_id": "s", + }, + { + "version": 1, + "command": "start", + "event_seq": True, + "request_id": "r", + "capture_id": "../escape", + "teleop_session_id": "s", + }, + ] + for value in bad_objects: + with self.subTest(value=value): + with self.assertRaises(ProtocolError): + parse_control_message(json.dumps(value)) + + def test_required_topic_preflight_demands_a_publisher_endpoint(self) -> None: + endpoints = { + "/robot_state": [object()], + # A discovered topic with no publisher is intentionally missing. + "/encoder_identical_joint": [], + "/hric/robot/rl_state": [object(), object()], + } + calls: list[str] = [] + + def lookup(topic: str) -> list[object]: + calls.append(topic) + return endpoints[topic] + + required = tuple(endpoints) + self.assertEqual( + topics_without_publishers(required, lookup), + ("/encoder_identical_joint",), + ) + self.assertEqual(calls, list(required)) + + +class RecorderControlProtocolTest(unittest.TestCase): + def setUp(self) -> None: + self.recorder = FakeRecorder() + self.protocol = RecorderControlProtocol( + self.recorder, + heartbeat_timeout_seconds=3.0, + monotonic=lambda: 0.0, + unix_time=lambda: 1234.5, + ) + + def test_start_is_idempotent_and_capture_id_is_preserved(self) -> None: + message = control("start", 1, "start-request") + first = self.protocol.handle_json(message, now=10.0) + second = self.protocol.handle_json(message, now=10.2) + self.assertEqual(len(self.recorder.starts), 1) + self.assertTrue(first["ack_accepted"]) + self.assertEqual(first["ack_code"], "start_accepted") + self.assertEqual(second["ack_code"], "start_accepted") + self.assertEqual(second["statistics"]["duplicates"], 1) + self.assertEqual(second["capture_id"], "capture-001") + + def test_request_id_reuse_and_stale_sequences_are_rejected(self) -> None: + self.protocol.handle_json(control("start", 5, "same"), now=1.0) + reused = self.protocol.handle_json( + control("heartbeat", 6, "same"), now=1.1 + ) + self.assertFalse(reused["ack_accepted"]) + self.assertEqual(reused["ack_code"], "request_id_reused") + + stale = self.protocol.handle_json( + control("heartbeat", 4, "different"), now=1.2 + ) + self.assertFalse(stale["ack_accepted"]) + self.assertEqual(stale["ack_code"], "stale_event_seq") + + def test_matching_heartbeat_and_stop_are_immediate(self) -> None: + self.protocol.handle_json(control("start", 1, "start"), now=5.0) + self.recorder.snapshot["state"] = "recording" + heartbeat = self.protocol.handle_json( + control("heartbeat", 2, "heartbeat"), now=6.0 + ) + self.assertTrue(heartbeat["ack_accepted"]) + self.assertEqual(heartbeat["ack_code"], "heartbeat_accepted") + stopped = self.protocol.handle_json( + control("stop", 3, "stop", reason="teleop_disarmed"), now=6.1 + ) + self.assertTrue(stopped["ack_accepted"]) + self.assertEqual(self.recorder.stops, ["teleop_disarmed"]) + + def test_wrong_capture_cannot_heartbeat_or_stop_active_episode(self) -> None: + self.protocol.handle_json(control("start", 1, "start"), now=5.0) + self.recorder.snapshot["state"] = "recording" + bad_heartbeat = self.protocol.handle_json( + control( + "heartbeat", + 2, + "bad-heartbeat", + capture_id="another-capture", + ), + now=5.5, + ) + bad_stop = self.protocol.handle_json( + control("stop", 3, "bad-stop", capture_id="another-capture"), + now=5.6, + ) + self.assertFalse(bad_heartbeat["ack_accepted"]) + self.assertFalse(bad_stop["ack_accepted"]) + self.assertEqual(self.recorder.stops, []) + + def test_missing_bridge_heartbeat_stops_recording_once(self) -> None: + self.protocol.handle_json(control("start", 1, "start"), now=10.0) + self.recorder.snapshot["state"] = "recording" + before = self.protocol.poll(now=13.0) + self.assertEqual(before["statistics"]["watchdog_stops"], 0) + after = self.protocol.poll(now=13.001) + again = self.protocol.poll(now=30.0) + self.assertEqual(self.recorder.stops, ["bridge_heartbeat_timeout"]) + self.assertEqual(after["statistics"]["watchdog_stops"], 1) + self.assertEqual(again["statistics"]["watchdog_stops"], 1) + + def test_manager_completion_maps_to_ready_or_failed_status(self) -> None: + self.protocol.handle_json(control("start", 1, "start"), now=1.0) + self.recorder.snapshot.update( + { + "state": "idle", + "episode_id": None, + "session_id": None, + "last_result": "ready", + } + ) + ready = self.protocol.poll(now=2.0) + self.assertEqual(ready["state"], "ready") + self.assertEqual(ready["capture_id"], "capture-001") + self.recorder.snapshot.update( + {"last_result": "failed", "last_error": "bag invalid"} + ) + failed = self.protocol.poll(now=2.1) + self.assertEqual(failed["state"], "failed") + self.assertEqual(failed["last_error"], "bag invalid") + + def test_preflight_failure_is_status_only(self) -> None: + protocol = RecorderControlProtocol( + self.recorder, + start_preflight=lambda: (False, "missing /robot_state"), + ) + status = protocol.handle_json(control("start", 1, "start"), now=1.0) + self.assertFalse(status["ack_accepted"]) + self.assertEqual(status["ack_code"], "preflight_failed") + self.assertEqual(status["last_error"], "missing /robot_state") + self.assertEqual(self.recorder.starts, []) + + def test_invalid_json_never_reaches_recorder(self) -> None: + status = self.protocol.handle_json("not-json", now=1.0) + self.assertFalse(status["ack_accepted"]) + self.assertEqual(status["ack_code"], "invalid_json") + self.assertEqual(self.recorder.starts, []) + self.assertEqual(self.recorder.stops, []) + + +if __name__ == "__main__": + unittest.main() diff --git a/tg3_local_teleop/test_session_gate.py b/tg3_local_teleop/test_session_gate.py index 917db5f..de068d6 100644 --- a/tg3_local_teleop/test_session_gate.py +++ b/tg3_local_teleop/test_session_gate.py @@ -30,6 +30,7 @@ for package, names in { "geometry_msgs.msg": ("TwistStamped",), "ros2_bridge_msgs.msg": ("ArmStatus",), "sensor_msgs.msg": ("JointState",), + "std_msgs.msg": ("String",), "std_srvs.srv": ("Trigger",), }.items(): install_module(package, **{name: Dummy for name in names}) @@ -42,6 +43,7 @@ sys.modules[spec.name] = bridge_module spec.loader.exec_module(bridge_module) ArmSnapshot = bridge_module.ArmSnapshot LocalTeleopBridge = bridge_module.LocalTeleopBridge +RecordingToggleGate = bridge_module.RecordingToggleGate class NullLogger: @@ -321,6 +323,158 @@ class RobotSessionGateTest(unittest.TestCase): bridge._tick_locomotion(1.4, old_wrong_binding) self.assertEqual(published[-1], (0.0, 0.0)) + def test_l3_hold_reaches_nonblocking_capture_request(self) -> None: + bridge = LocalTeleopBridge.__new__(LocalTeleopBridge) + bridge.data_collection_enabled = True + bridge.data_collection_cfg = { + "button_input_timeout_s": 0.25, + "control_retry_seconds": 0.5, + "heartbeat_interval_seconds": 0.5, + } + bridge.data_collection_gate = RecordingToggleGate(1.0, 0.5) + bridge.data_collection_gate.new_session() + bridge.armed = True + bridge.data_pending_control = None + bridge.data_last_control_publish_at = 0.0 + bridge.data_recorder_status = {} + bridge.data_recorder_status_at = 0.0 + bridge.data_last_heartbeat_at = 0.0 + bridge.data_last_iarm_received_at = 0.0 + bridge.data_iarm_publisher = None + bridge.data_toggle_count = 0 + requested: list[str] = [] + bridge._request_data_capture = MethodType( + lambda _self, command, _reason: requested.append(command), bridge + ) + + def frame(now: float, pressed: int) -> object: + return ArmSnapshot( + {"button_joystick": {"left": pressed}}, received_at=now + ) + + bridge._tick_data_collection(10.0, frame(10.0, 0)) + bridge._tick_data_collection(10.5, frame(10.5, 0)) + bridge._tick_data_collection(10.6, frame(10.6, 1)) + bridge._tick_data_collection(11.61, frame(11.61, 1)) + self.assertEqual(requested, ["start"]) + self.assertEqual(bridge.data_toggle_count, 1) + + def test_rejected_recorder_start_clears_requested_active_state(self) -> None: + bridge = LocalTeleopBridge.__new__(LocalTeleopBridge) + bridge.data_collection_gate = RecordingToggleGate(1.0, 0.5) + bridge.data_collection_gate.new_session() + bridge.data_collection_gate.active = True + bridge.data_pending_control = { + "command": "start", + "request_id": "request-1", + "event_seq": 4, + } + bridge.data_recorder_status = {} + bridge.data_recorder_status_at = 0.0 + bridge.data_capture_id = "capture-1" + bridge.data_last_transition = "start_requested" + bridge.get_logger = MethodType(lambda _self: NullLogger(), bridge) + message = Dummy() + message.data = ( + '{"version":1,"state":"failed","capture_id":"capture-1",' + '"ack_request_id":"request-1","ack_event_seq":4,' + '"ack_accepted":false,"ack_code":"preflight_failed",' + '"last_error":"required topic missing"}' + ) + + bridge._on_data_recorder_status(message) + self.assertIsNone(bridge.data_pending_control) + self.assertFalse(bridge.data_collection_gate.active) + self.assertIn("start_rejected", bridge.data_last_transition) + + def test_supervisor_context_loss_clears_false_recording_state(self) -> None: + bridge = LocalTeleopBridge.__new__(LocalTeleopBridge) + bridge.data_collection_gate = RecordingToggleGate(1.0, 0.5) + bridge.data_collection_gate.new_session() + bridge.data_collection_gate.active = True + bridge.data_pending_control = None + bridge.data_recorder_status = {} + bridge.data_recorder_status_at = 0.0 + bridge.data_capture_id = "capture-old" + bridge.data_last_transition = "recording" + bridge.get_logger = MethodType(lambda _self: NullLogger(), bridge) + message = Dummy() + message.data = ( + '{"version":1,"state":"idle","capture_id":null,' + '"ack_request_id":null,"ack_event_seq":null,' + '"ack_accepted":null,"last_error":""}' + ) + + bridge._on_data_recorder_status(message) + self.assertFalse(bridge.data_collection_gate.active) + self.assertEqual( + bridge.data_last_transition, + "recorder context was lost or replaced", + ) + + def test_missing_supervisor_ack_cancels_false_start(self) -> None: + bridge = LocalTeleopBridge.__new__(LocalTeleopBridge) + bridge.data_collection_enabled = True + bridge.data_collection_cfg = { + "control_retry_seconds": 0.5, + "ack_timeout_seconds": 5.0, + "heartbeat_interval_seconds": 0.5, + "status_stale_seconds": 4.0, + } + bridge.data_collection_gate = RecordingToggleGate(1.0, 0.5) + bridge.data_collection_gate.new_session() + bridge.data_collection_gate.active = True + bridge.armed = False + bridge.data_pending_control = {"command": "start"} + bridge.data_pending_control_since = 10.0 + bridge.data_last_control_publish_at = 10.0 + bridge.data_control_publisher = None + bridge.data_recorder_status = {} + bridge.data_recorder_status_at = 0.0 + bridge.data_last_heartbeat_at = 0.0 + bridge.data_last_iarm_received_at = 0.0 + bridge.data_iarm_publisher = None + bridge.data_last_transition = "start_requested" + bridge.get_logger = MethodType(lambda _self: NullLogger(), bridge) + + bridge._tick_data_collection(15.0, None) + self.assertIsNone(bridge.data_pending_control) + self.assertFalse(bridge.data_collection_gate.active) + self.assertEqual(bridge.data_last_transition, "start_ack_timeout") + + def test_stale_recorder_status_requests_capture_stop(self) -> None: + bridge = LocalTeleopBridge.__new__(LocalTeleopBridge) + bridge.data_collection_enabled = True + bridge.data_collection_cfg = { + "control_retry_seconds": 0.5, + "ack_timeout_seconds": 5.0, + "heartbeat_interval_seconds": 0.5, + "status_stale_seconds": 4.0, + } + bridge.data_collection_gate = RecordingToggleGate(1.0, 0.5) + bridge.data_collection_gate.new_session() + bridge.data_collection_gate.active = True + bridge.armed = False + bridge.data_pending_control = None + bridge.data_pending_control_since = 0.0 + bridge.data_last_control_publish_at = 0.0 + bridge.data_recorder_status = {"state": "recording"} + bridge.data_recorder_status_at = 10.0 + bridge.data_last_heartbeat_at = 10.0 + bridge.data_last_iarm_received_at = 0.0 + bridge.data_iarm_publisher = None + bridge.data_last_transition = "recording" + bridge.get_logger = MethodType(lambda _self: NullLogger(), bridge) + requested: list[str] = [] + bridge._request_data_capture = MethodType( + lambda _self, command, _reason: requested.append(command), bridge + ) + + bridge._tick_data_collection(14.0, None) + self.assertFalse(bridge.data_collection_gate.active) + self.assertEqual(requested, ["stop"]) + self.assertEqual(bridge.data_last_transition, "recorder_status_stale") + if __name__ == "__main__": unittest.main() diff --git a/tg3_local_teleop/tg3-data-recorder.service b/tg3_local_teleop/tg3-data-recorder.service new file mode 100644 index 0000000..a01c1b5 --- /dev/null +++ b/tg3_local_teleop/tg3-data-recorder.service @@ -0,0 +1,23 @@ +[Unit] +Description=TG3 independent teleoperation MCAP recorder supervisor +After=network-online.target +Wants=network-online.target +Before=tg3-local-teleop.service + +[Service] +Type=simple +WorkingDirectory=/home/nvidia/tg3_local_teleop +UMask=0027 +ExecStartPre=/home/nvidia/tg3_local_teleop/wait_ros_ready.sh +ExecStart=/home/nvidia/tg3_local_teleop/run_data_recorder.sh +Restart=on-failure +RestartSec=2 +KillSignal=SIGINT +# Signal only the Python supervisor first; it owns clean SIGINT/finalization +# of rosbag2. The whole cgroup is still killed if the generous deadline ends. +KillMode=mixed +TimeoutStartSec=150 +TimeoutStopSec=150 + +[Install] +WantedBy=default.target diff --git a/tg3_local_teleop/tg3-local-teleop.service b/tg3_local_teleop/tg3-local-teleop.service index 70281eb..aa6512d 100644 --- a/tg3_local_teleop/tg3-local-teleop.service +++ b/tg3_local_teleop/tg3-local-teleop.service @@ -1,7 +1,7 @@ [Unit] Description=TG3 local TS1P dual-arm bridge (no cloud pairing) -After=network-online.target -Wants=network-online.target +After=network-online.target tg3-data-recorder.service +Wants=network-online.target tg3-data-recorder.service [Service] Type=simple diff --git a/tg3_local_teleop/tg3_local_teleop.py b/tg3_local_teleop/tg3_local_teleop.py index 04b9577..35c1f00 100755 --- a/tg3_local_teleop/tg3_local_teleop.py +++ b/tg3_local_teleop/tg3_local_teleop.py @@ -16,6 +16,7 @@ import struct import threading import time import tomllib +import uuid from dataclasses import dataclass from pathlib import Path from typing import Any @@ -28,8 +29,10 @@ from geometry_msgs.msg import TwistStamped from rclpy.node import Node from ros2_bridge_msgs.msg import ArmStatus from sensor_msgs.msg import JointState +from std_msgs.msg import String from std_srvs.srv import Trigger +from data_collection import RecordingToggleGate, left_joystick_pressed from gesture_toggle import ( GestureToggle, normalized_pose_to_positions, @@ -426,6 +429,32 @@ class LocalTeleopBridge(Node): self.locomotion_enabled = bool( self.locomotion_cfg.get("enabled", False) ) + self.data_collection_cfg = config.get("data_collection", {}) + self.data_collection_enabled = bool( + self.data_collection_cfg.get("enabled", False) + ) + self.data_collection_gate = RecordingToggleGate( + hold_seconds=float( + self.data_collection_cfg.get("button_hold_seconds", 1.0) + ), + release_seconds=float( + self.data_collection_cfg.get("button_release_seconds", 0.5) + ), + ) + self.data_control_publisher = None + self.data_iarm_publisher = None + self.data_recorder_status: dict[str, Any] = {} + self.data_recorder_status_at = 0.0 + self.data_pending_control: dict[str, Any] | None = None + self.data_pending_control_since = 0.0 + self.data_last_control_publish_at = 0.0 + self.data_last_heartbeat_at = 0.0 + self.data_last_iarm_received_at = 0.0 + self.data_capture_id: str | None = None + self.data_capture_session_id: str | None = None + self.data_event_seq = 0 + self.data_toggle_count = 0 + self.data_last_transition = "initialized; waiting for an armed session" self.hand_publishers: dict[str, Any] = {} self.robot_hand_positions: dict[str, list[int] | None] = { side: None for side in HAND_SIDES @@ -492,6 +521,23 @@ class LocalTeleopBridge(Node): self.create_service( Trigger, ros_cfg["cancel_home_service"], self._on_cancel_home_request ) + if getattr(self, "data_collection_enabled", False): + self.data_control_publisher = self.create_publisher( + String, + str(self.data_collection_cfg["control_topic"]), + 10, + ) + self.data_iarm_publisher = self.create_publisher( + String, + str(self.data_collection_cfg["iarm_frame_topic"]), + 10, + ) + self.create_subscription( + String, + str(self.data_collection_cfg["status_topic"]), + self._on_data_recorder_status, + 10, + ) self.source = LatestArmData(net_cfg) self.source.start() @@ -526,7 +572,8 @@ class LocalTeleopBridge(Node): self.get_logger().info( f"local bridge started in {mode}; source={self.source.description}; " f"target={ros_cfg['command_topic']}; brainco_hands={self.hands_enabled}; " - f"locomotion={self.locomotion_enabled}" + f"locomotion={self.locomotion_enabled}; " + f"data_collection={self.data_collection_enabled}" ) def close(self) -> None: @@ -558,6 +605,242 @@ class LocalTeleopBridge(Node): self.robot_hand_states[side] = states self.robot_hand_at[side] = time.monotonic() + def _on_data_recorder_status(self, msg: String) -> None: + """Consume recorder acknowledgements without affecting robot control.""" + + try: + status = json.loads(msg.data) + if not isinstance(status, dict) or status.get("version") != 1: + raise ValueError("unsupported recorder status") + except (TypeError, ValueError, json.JSONDecodeError) as exc: + self.get_logger().warning(f"ignored invalid data-recorder status: {exc}") + return + + self.data_recorder_status = status + self.data_recorder_status_at = time.monotonic() + pending = self.data_pending_control + if ( + pending is not None + and status.get("ack_request_id") == pending.get("request_id") + and status.get("ack_event_seq") == pending.get("event_seq") + ): + accepted = status.get("ack_accepted") + if type(accepted) is bool: + self.data_pending_control = None + self.data_pending_control_since = 0.0 + if not accepted: + detail = str( + status.get("last_error") + or status.get("ack_code") + or "recorder rejected request" + ) + if pending.get("command") == "start": + self.data_collection_gate.force_inactive( + f"recorder_rejected: {detail}" + ) + self.data_last_transition = ( + f"{pending.get('command')}_rejected: {detail}" + ) + self.get_logger().error( + "data recorder rejected %s: %s" + % (pending.get("command"), detail) + ) + + state = status.get("state") + if ( + state in ("ready", "failed") + and status.get("capture_id") == self.data_capture_id + and self.data_collection_gate.active + ): + self.data_collection_gate.force_inactive( + f"recorder_{state}: {status.get('last_error', '')}".rstrip() + ) + self.data_last_transition = self.data_collection_gate.last_transition + elif ( + self.data_pending_control is None + and self.data_collection_gate.active + and state in ("idle", "ready", "failed", "recording", "stopping") + and status.get("capture_id") != self.data_capture_id + ): + # A restarted supervisor has no in-memory context for the old bag. + # Reconcile the UI gate instead of displaying a false recording + # state forever; the supervisor owns cleanup of its old cgroup and + # active/failed directory. + self.data_collection_gate.force_inactive( + "recorder context was lost or replaced" + ) + self.data_last_transition = self.data_collection_gate.last_transition + self.get_logger().error( + "data recorder no longer owns the requested capture; " + "release L3 before starting a new episode" + ) + + def _publish_data_control(self, payload: dict[str, Any]) -> None: + if self.data_control_publisher is None: + return + try: + message = String() + message.data = json.dumps( + payload, ensure_ascii=False, separators=(",", ":") + ) + self.data_control_publisher.publish(message) + self.data_last_control_publish_at = time.monotonic() + except Exception as exc: # Data collection must never stop robot control. + self.get_logger().error(f"cannot publish data-recorder control: {exc}") + + def _request_data_capture(self, command: str, reason: str) -> None: + if not self.data_collection_enabled: + return + if command == "start": + self.data_capture_id = ( + time.strftime("%Y%m%dT%H%M%SZ", time.gmtime()) + + "_TG3_" + + uuid.uuid4().hex[:8] + ) + if self.data_capture_id is None: + return + if self.data_capture_session_id is None: + self.data_capture_session_id = self.active_session_id or ( + "direct_" + uuid.uuid4().hex + ) + + self.data_event_seq += 1 + payload: dict[str, Any] = { + "version": 1, + "command": command, + "event_seq": self.data_event_seq, + "request_id": uuid.uuid4().hex, + "capture_id": self.data_capture_id, + "teleop_session_id": self.data_capture_session_id, + "reason": reason, + "sent_unix_s": time.time(), + } + self.data_pending_control = payload + self.data_pending_control_since = time.monotonic() + self._publish_data_control(payload) + self.data_last_transition = f"{command}_requested: {reason}" + + def _publish_data_heartbeat(self) -> None: + if self.data_capture_id is None or self.data_capture_session_id is None: + return + self.data_event_seq += 1 + payload = { + "version": 1, + "command": "heartbeat", + "event_seq": self.data_event_seq, + "request_id": uuid.uuid4().hex, + "capture_id": self.data_capture_id, + "teleop_session_id": self.data_capture_session_id, + "sent_unix_s": time.time(), + } + self._publish_data_control(payload) + self.data_last_heartbeat_at = time.monotonic() + + def _tick_data_collection( + self, now: float, sample: ArmSnapshot | None + ) -> None: + if not self.data_collection_enabled: + return + + if self.armed: + input_timeout = float( + self.data_collection_cfg.get("button_input_timeout_s", 0.25) + ) + input_healthy = ( + sample is not None and now - sample.received_at <= input_timeout + ) + pressed = ( + None if sample is None else left_joystick_pressed(sample.data) + ) + action = self.data_collection_gate.update( + now, + input_healthy=input_healthy, + pressed=pressed, + ) + if action is not None: + self.data_toggle_count += 1 + self._request_data_capture( + action, + "left joystick press held for " + f"{self.data_collection_gate.hold_seconds:.1f}s", + ) + + retry_s = float(self.data_collection_cfg.get("control_retry_seconds", 0.5)) + if ( + self.data_pending_control is not None + and now - self.data_last_control_publish_at >= retry_s + ): + self._publish_data_control(self.data_pending_control) + + ack_timeout_s = float( + self.data_collection_cfg.get("ack_timeout_seconds", 5.0) + ) + if ( + self.data_pending_control is not None + and self.data_pending_control_since > 0.0 + and now - self.data_pending_control_since >= ack_timeout_s + ): + expired_command = str(self.data_pending_control.get("command")) + self.data_pending_control = None + self.data_pending_control_since = 0.0 + if expired_command == "start": + self.data_collection_gate.force_inactive( + "recorder START acknowledgement timed out" + ) + self.data_last_transition = f"{expired_command}_ack_timeout" + self.get_logger().error( + f"data recorder {expired_command} acknowledgement timed out" + ) + + status_stale_s = float( + self.data_collection_cfg.get("status_stale_seconds", 4.0) + ) + if ( + self.data_collection_gate.active + and self.data_pending_control is None + and self.data_recorder_status_at > 0.0 + and now - self.data_recorder_status_at >= status_stale_s + ): + self.data_collection_gate.force_inactive( + "recorder status heartbeat became stale" + ) + self.data_last_transition = "recorder_status_stale" + self._request_data_capture( + "stop", "recorder status heartbeat became stale" + ) + self.get_logger().error( + "data recorder status became stale; capture stop requested" + ) + + recorder_state = self.data_recorder_status.get("state") + heartbeat_s = float( + self.data_collection_cfg.get("heartbeat_interval_seconds", 0.5) + ) + if ( + self.data_collection_gate.active + and self.data_pending_control is None + and recorder_state in ("starting", "recording") + and now - self.data_last_heartbeat_at >= heartbeat_s + ): + self._publish_data_heartbeat() + + if ( + self.data_collection_gate.active + and sample is not None + and sample.received_at != self.data_last_iarm_received_at + and self.data_iarm_publisher is not None + ): + try: + message = String() + # Preserve the complete xTELE/OmniSocket application frame. + message.data = json.dumps( + sample.data, ensure_ascii=False, separators=(",", ":") + ) + self.data_iarm_publisher.publish(message) + self.data_last_iarm_received_at = sample.received_at + except Exception as exc: + self.get_logger().error(f"cannot publish xTELE capture frame: {exc}") + @staticmethod def _hand_message_signature(msg: SetMotorMulti) -> tuple[Any, ...]: return ( @@ -690,6 +973,11 @@ class LocalTeleopBridge(Node): success=False, ) + # The recorder toggle is deliberately downstream of every disarm + # decision, so a matching STOP or safety teardown always wins over an + # L3 press observed in the same control tick. + self._tick_data_collection(now, sample) + if self.right_point_gesture_enabled: gesture_toggled = self.right_point_gesture.update( now, @@ -937,6 +1225,10 @@ class LocalTeleopBridge(Node): self.active_session_id = session_id if self.right_point_gesture_enabled: self.right_point_gesture.new_session() + if getattr(self, "data_collection_enabled", False): + self.data_capture_session_id = session_id + self.data_collection_gate.new_session() + self.data_last_transition = "new teleoperation session" self.armed = True # Start both slew limiters at measured robot feedback, never at a # potentially distant first network target. @@ -1002,6 +1294,10 @@ class LocalTeleopBridge(Node): return if self.right_point_gesture_enabled: self.right_point_gesture.new_session() + if getattr(self, "data_collection_enabled", False): + self.data_capture_session_id = "direct_" + uuid.uuid4().hex + self.data_collection_gate.new_session() + self.data_last_transition = "new direct-LAN teleoperation session" self.armed = True # Start the slew limiter at measured robot feedback. Using None here # would make the first armed frame jump directly to the TS1P target. @@ -1020,6 +1316,13 @@ class LocalTeleopBridge(Node): def _disarm(self, reason: str) -> None: was_armed = self.armed + if getattr(self, "data_collection_enabled", False): + action = self.data_collection_gate.end_session(reason) + self.data_last_transition = self.data_collection_gate.last_transition + if action == "stop": + # This is a non-blocking ROS request. It never delays robot + # disarm, locomotion zeroing, or the existing Home sequence. + self._request_data_capture("stop", reason) self._stop_locomotion(reason) if self.right_point_gesture_enabled: # Clearing the logical override must not publish a hand target. @@ -1550,6 +1853,46 @@ class LocalTeleopBridge(Node): "custom_hand_feedback_policy": ( "startup_qualifies_teleop; runtime_gap_pauses_hands_only" ), + "data_collection_enabled": self.data_collection_enabled, + "data_collection_binding": ( + "left joystick press (L3), hold " + f"{self.data_collection_gate.hold_seconds:.1f}s toggle; " + f"release {self.data_collection_gate.release_seconds:.1f}s" + ), + "data_collection_gate_state": self.data_collection_gate.state, + "data_collection_requested_active": self.data_collection_gate.active, + "data_collection_requires_release": ( + self.data_collection_gate.require_release + ), + "data_collection_button_pressed": ( + self.data_collection_gate.button_pressed + ), + "data_collection_hold_s": ( + 0.0 + if self.data_collection_gate.hold_started_at is None + else round(now - self.data_collection_gate.hold_started_at, 2) + ), + "data_collection_capture_id": self.data_capture_id, + "data_collection_session_id": self.data_capture_session_id, + "data_collection_toggle_count": self.data_toggle_count, + "data_collection_last_transition": self.data_last_transition, + "data_collection_pending_command": ( + None + if self.data_pending_control is None + else self.data_pending_control.get("command") + ), + "data_collection_pending_age_s": ( + None + if self.data_pending_control is None + or self.data_pending_control_since == 0.0 + else round(now - self.data_pending_control_since, 3) + ), + "data_recorder_status": self.data_recorder_status, + "data_recorder_status_age_s": ( + None + if self.data_recorder_status_at == 0.0 + else round(now - self.data_recorder_status_at, 3) + ), "runtime_hand_output_ready": self.hand_output_ready, "runtime_hand_output_reasons": self.runtime_hand_output_reasons, "safety_ready": not reasons, diff --git a/tg3_omnisocket_transport/tg3-omnisocket-sender.service b/tg3_omnisocket_transport/tg3-omnisocket-sender.service index 62320a9..0cbe7dc 100644 --- a/tg3_omnisocket_transport/tg3-omnisocket-sender.service +++ b/tg3_omnisocket_transport/tg3-omnisocket-sender.service @@ -7,7 +7,7 @@ Wants=network-online.target Type=simple WorkingDirectory=/home/eai/tg3_omnisocket_transport Environment=PYTHONPATH=/home/eai/OmniSocketGo/python -ExecStart=/usr/bin/python3 /home/eai/tg3_omnisocket_transport/omnisocket_xtele_sender.py --server 175.178.116.187:14049 --peer-id tg3-009027fa8190-iarm --target-peer tg3-009027fa8190-robot --zmq-endpoint tcp://127.0.0.1:5003 --cmd-zmq-endpoint tcp://127.0.0.1:5001 --cmd-max-age-s 0.25 --source-timeout-s 0.25 --start-stop-hold-s 3.0 --combo-release-s 0.5 --start-marker-frames 500 --max-feedback-age-ms 500 --max-pending-frames 100 --status-file /home/eai/tg3_omnisocket_transport/status.json +ExecStart=/usr/bin/python3 /home/eai/tg3_omnisocket_transport/omnisocket_xtele_sender.py --server 127.0.0.1:14049 --peer-id tg3-009027fa8190-iarm --target-peer tg3-009027fa8190-robot --zmq-endpoint tcp://127.0.0.1:5003 --cmd-zmq-endpoint tcp://127.0.0.1:5001 --cmd-max-age-s 0.25 --source-timeout-s 0.25 --start-stop-hold-s 3.0 --combo-release-s 0.5 --start-marker-frames 500 --max-feedback-age-ms 500 --max-pending-frames 100 --status-file /home/eai/tg3_omnisocket_transport/status.json Restart=always RestartSec=1 KillSignal=SIGINT diff --git a/verify.sh b/verify.sh index 5c2b01f..ed4eefa 100755 --- a/verify.sh +++ b/verify.sh @@ -7,10 +7,13 @@ expected_omnisocket_commit="de3f5c96779dbe1571c10feb22fc7f2331b6b222" for executable in \ "$repo_dir/tg3_omnisocket_transport/omnisocket_xtele_sender.py" \ "$repo_dir/tg3_local_teleop/tg3_local_teleop.py" \ + "$repo_dir/tg3_local_teleop/data_recorder_node.py" \ "$repo_dir/tg3_local_teleop/run.sh" \ + "$repo_dir/tg3_local_teleop/run_data_recorder.sh" \ "$repo_dir/tg3_local_teleop/wait_ros_ready.sh" \ "$repo_dir/tg3_local_teleop/home.sh" \ - "$repo_dir/tg3_local_teleop/status.sh"; do + "$repo_dir/tg3_local_teleop/status.sh" \ + "$repo_dir/tg3_data_collection/data_get_sync.py"; do if [[ ! -x "$executable" ]]; then echo "Required executable bit is missing: $executable" >&2 exit 1 @@ -20,11 +23,18 @@ done python3 -m py_compile \ "$repo_dir/tg3_omnisocket_transport/omnisocket_xtele_sender.py" \ "$repo_dir/tg3_local_teleop/tg3_local_teleop.py" \ - "$repo_dir/tg3_local_teleop/gesture_toggle.py" + "$repo_dir/tg3_local_teleop/gesture_toggle.py" \ + "$repo_dir/tg3_local_teleop/data_collection.py" \ + "$repo_dir/tg3_local_teleop/data_recorder_protocol.py" \ + "$repo_dir/tg3_local_teleop/data_recorder_node.py" \ + "$repo_dir/tg3_data_collection/data_get_sync.py" python3 "$repo_dir/tg3_omnisocket_transport/test_session_gate.py" python3 "$repo_dir/tg3_local_teleop/test_session_gate.py" python3 "$repo_dir/tg3_local_teleop/test_gesture_toggle.py" python3 "$repo_dir/tg3_local_teleop/test_idle_session_refresh.py" +python3 "$repo_dir/tg3_local_teleop/test_data_collection.py" +python3 "$repo_dir/tg3_local_teleop/test_data_recorder_protocol.py" +python3 "$repo_dir/tg3_data_collection/test_data_get_sync.py" python3 - "$repo_dir/tg3_local_teleop/config.toml" <<'PY' import sys import tomllib