feat: add session-gated TG3 data collection
This commit is contained in:
363
tg3_data_collection/data_get_sync.py
Executable file
363
tg3_data_collection/data_get_sync.py
Executable file
@@ -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())
|
||||
Reference in New Issue
Block a user