Files
TG3/tg3_data_collection/data_get_sync.py

366 lines
13 KiB
Python
Executable File

#!/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"if test -d {root}; then "
f"find {root} -mindepth 1 -maxdepth 1 -type d -printf '%f\\n'; "
"fi"
)
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())