Files
TG3/tg3_data_collection/test_data_get_sync.py

600 lines
26 KiB
Python
Executable File

#!/usr/bin/env python3
from __future__ import annotations
import hashlib
import json
import os
import shutil
import subprocess
import tempfile
import unittest
from pathlib import Path
from unittest import mock
from data_get_sync import (
DataGetSync,
FIXED_REMOTE_READY,
ManifestDocument,
VERIFIED_RECEIPT_NAME,
fsync_episode_tree,
parse_manifest,
safe_episode_name,
safe_relative_path,
validate_episode_dir,
)
class LocalRsyncRunner:
def __init__(self, remote_root: Path) -> None:
self.remote_root = remote_root
self.mode = "normal"
self.rsync_calls = 0
def run(
self, command: list[str], *, timeout: float
) -> subprocess.CompletedProcess[str]:
del timeout
if not command or command[0] != "rsync":
return subprocess.CompletedProcess(command, 1, "", "unexpected command")
self.rsync_calls += 1
remote_source = command[-2].rstrip("/")
episode = remote_source.rsplit("/", 1)[-1]
destination = Path(command[-1])
shutil.copytree(
self.remote_root / episode,
destination,
dirs_exist_ok=True,
symlinks=True,
)
if self.mode == "truncate":
(destination / "bag" / "bag_0.mcap").write_bytes(b"truncated")
elif self.mode == "manifest_changed":
manifest = destination / "manifest.json"
manifest.write_bytes(manifest.read_bytes() + b"\n")
elif self.mode == "manifest_symlink":
manifest = destination / "manifest.json"
manifest.unlink()
manifest.symlink_to("bag/metadata.yaml")
return subprocess.CompletedProcess(command, 0, "", "")
class LocalSync(DataGetSync):
def __init__(
self,
*,
remote_root: Path,
destination: Path,
runner: LocalRsyncRunner,
delete: bool = True,
remote: str = "nvidia@test",
) -> None:
super().__init__(
remote=remote,
remote_ready=FIXED_REMOTE_READY,
destination=destination,
status_file=destination / "sync_status.json",
runner=runner,
reserve_bytes=0,
delete_remote_after_sync=delete,
remote_delete_helper="/fixed/delete_ready_episode.py",
)
self.remote_root = remote_root
self.delete_calls: list[tuple[str, str]] = []
self.delete_outcomes: list[str] = []
def list_remote_episodes(self) -> list[str]:
return sorted(
path.name
for path in self.remote_root.iterdir()
if path.is_dir() and not path.is_symlink()
)
def get_remote_manifest(self, episode: str) -> ManifestDocument:
raw = (self.remote_root / episode / "manifest.json").read_bytes()
return ManifestDocument.from_bytes(raw, episode)
def _invoke_delete_helper(
self, episode: str, manifest_sha256: str
) -> dict[str, object]:
self.delete_calls.append((episode, manifest_sha256))
outcome = self.delete_outcomes.pop(0) if self.delete_outcomes else "success"
target = self.remote_root / episode
if outcome == "failure":
raise RuntimeError("injected delete failure")
if outcome == "ack_lost_after_delete":
if target.exists():
shutil.rmtree(target)
raise RuntimeError("injected lost acknowledgement")
if target.exists():
shutil.rmtree(target)
state = "deleted"
else:
state = "already_absent"
return {
"state": state,
"episode_id": episode,
"manifest_sha256": manifest_sha256,
}
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_fsync_episode_tree_rejects_intermediate_symlink(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
episode = root / "episode"
episode.mkdir()
outside = root / "outside"
outside.mkdir()
(outside / "file").write_bytes(b"sentinel")
(episode / "linked").symlink_to(outside, target_is_directory=True)
with self.assertRaises(ValueError):
fsync_episode_tree(episode)
def _local_sync(
self, root: Path, name: str = "episode_20260810T120000000_deadbeef"
) -> tuple[LocalSync, LocalRsyncRunner, Path, Path]:
remote = root / "remote_ready"
destination = root / "Data_Get"
remote.mkdir()
destination.mkdir()
self._episode(remote, name)
runner = LocalRsyncRunner(remote)
syncer = LocalSync(
remote_root=remote,
destination=destination,
runner=runner,
)
return syncer, runner, remote, destination
def test_copy_is_verified_receipted_then_remote_is_deleted(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
syncer, _runner, remote, destination = self._local_sync(
Path(temporary), name
)
self.assertEqual(syncer.run_once(), 1)
final = destination / name
receipt = json.loads((final / VERIFIED_RECEIPT_NAME).read_text())
self.assertEqual(receipt["state"], "VERIFIED")
self.assertEqual(receipt["remote_identity"], syncer.remote_identity)
self.assertEqual(receipt["verified_files"], 2)
self.assertGreater(receipt["verified_bytes"], 0)
self.assertEqual(receipt["remote_delete"]["state"], "deleted")
self.assertFalse((remote / name).exists())
self.assertEqual(syncer.delete_count, 1)
status = json.loads((destination / "sync_status.json").read_text())
self.assertEqual(status["state"], "idle")
self.assertEqual(status["pending_remote_cleanup"], 0)
def test_truncated_copy_never_calls_delete(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
syncer, runner, remote, _destination = self._local_sync(Path(temporary))
runner.mode = "truncate"
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(syncer.delete_calls, [])
self.assertEqual(len(list(remote.iterdir())), 1)
def test_manifest_change_during_copy_never_calls_delete(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
syncer, runner, remote, _destination = self._local_sync(Path(temporary))
runner.mode = "manifest_changed"
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(syncer.delete_calls, [])
self.assertEqual(len(list(remote.iterdir())), 1)
def test_copied_manifest_symlink_never_calls_delete(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
syncer, runner, remote, _destination = self._local_sync(Path(temporary))
runner.mode = "manifest_symlink"
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(syncer.delete_calls, [])
self.assertEqual(len(list(remote.iterdir())), 1)
def test_delete_failure_is_pending_and_retried(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
syncer, _runner, remote, destination = self._local_sync(
Path(temporary), name
)
syncer.delete_outcomes = ["failure", "success"]
self.assertEqual(syncer.run_once(), 1)
receipt_path = destination / name / VERIFIED_RECEIPT_NAME
receipt = json.loads(receipt_path.read_text())
self.assertEqual(receipt["remote_delete"]["state"], "pending")
self.assertTrue((remote / name).exists())
status = json.loads((destination / "sync_status.json").read_text())
self.assertEqual(status["state"], "delete_pending")
self.assertEqual(status["pending_remote_cleanup"], 1)
# Persistent backoff prevents a 2-second poll loop from repeatedly
# hashing a large MCAP after a network/helper failure.
self.assertEqual(syncer.run_once(), 0)
self.assertEqual(len(syncer.delete_calls), 1)
receipt = json.loads(receipt_path.read_text())
receipt["remote_delete"]["next_retry_unix_s"] = 0.0
receipt_path.write_text(json.dumps(receipt), encoding="utf-8")
self.assertEqual(syncer.run_once(), 0)
receipt = json.loads(receipt_path.read_text())
self.assertEqual(receipt["remote_delete"]["state"], "deleted")
self.assertFalse((remote / name).exists())
self.assertEqual(len(syncer.delete_calls), 2)
def test_lost_delete_ack_is_idempotent_after_process_restart(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
root = Path(temporary)
first, runner, remote, destination = self._local_sync(root, name)
first.delete_outcomes = ["ack_lost_after_delete"]
self.assertEqual(first.run_once(), 1)
self.assertFalse((remote / name).exists())
receipt_path = destination / name / VERIFIED_RECEIPT_NAME
self.assertEqual(
json.loads(receipt_path.read_text())["remote_delete"]["state"],
"pending",
)
restarted = LocalSync(
remote_root=remote,
destination=destination,
runner=runner,
)
with mock.patch(
"data_get_sync.validate_episode_dir",
wraps=validate_episode_dir,
) as deep_validator:
self.assertEqual(restarted.run_once(), 0)
self.assertEqual(deep_validator.call_count, 0)
self.assertEqual(restarted.delete_calls, [])
receipt = json.loads(receipt_path.read_text())
receipt["remote_delete"]["next_retry_unix_s"] = 0.0
receipt_path.write_text(json.dumps(receipt), encoding="utf-8")
self.assertEqual(restarted.run_once(), 0)
self.assertGreaterEqual(deep_validator.call_count, 1)
self.assertEqual(
json.loads(receipt_path.read_text())["remote_delete"]["state"],
"deleted",
)
self.assertEqual(restarted.delete_calls[0][0], name)
def test_retry_rehashes_and_refuses_delete_after_local_tamper(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
syncer, _runner, remote, destination = self._local_sync(
Path(temporary), name
)
syncer.delete_outcomes = ["failure", "success"]
self.assertEqual(syncer.run_once(), 1)
(destination / name / "bag" / "bag_0.mcap").write_bytes(b"tampered")
receipt_path = destination / name / VERIFIED_RECEIPT_NAME
receipt = json.loads(receipt_path.read_text())
receipt["remote_delete"]["next_retry_unix_s"] = 0.0
receipt_path.write_text(json.dumps(receipt), encoding="utf-8")
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(len(syncer.delete_calls), 1)
self.assertTrue((remote / name).exists())
def test_unlisted_stale_mcap_is_never_published_or_deleted(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
syncer, _runner, remote, destination = self._local_sync(
Path(temporary), name
)
(remote / name / "bag" / "stale_old_payload.mcap").write_bytes(
b"stale"
)
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertFalse((destination / name).exists())
self.assertEqual(syncer.delete_calls, [])
self.assertTrue((remote / name).exists())
def test_reused_hash_staging_cannot_smuggle_old_mcap(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
syncer, _runner, remote, destination = self._local_sync(
Path(temporary), name
)
document = syncer.get_remote_manifest(name)
staging = (
destination
/ ".incoming"
/ f"{name}.{document.sha256}.partial"
/ "bag"
)
staging.mkdir(parents=True)
(staging / "stale_old_payload.mcap").write_bytes(b"old interrupted bag")
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertFalse((destination / name).exists())
self.assertEqual(syncer.delete_calls, [])
self.assertTrue((remote / name).exists())
def test_deleted_receipt_from_old_ip_does_not_block_new_ip(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
root = Path(temporary)
first, runner, remote, destination = self._local_sync(root, name)
self.assertEqual(first.run_once(), 1)
receipt_path = destination / name / VERIFIED_RECEIPT_NAME
receipt = json.loads(receipt_path.read_text())
receipt["remote"] = "nvidia@old-ip"
receipt["remote_identity"] = (
f"nvidia@old-ip:{FIXED_REMOTE_READY}"
)
receipt_path.write_text(json.dumps(receipt), encoding="utf-8")
migrated = LocalSync(
remote_root=remote,
destination=destination,
runner=runner,
remote="nvidia@new-ip",
)
self.assertEqual(migrated.run_once(), 0)
status = json.loads((destination / "sync_status.json").read_text())
self.assertEqual(status["state"], "idle")
self.assertEqual(status["pending_remote_cleanup"], 0)
def test_bad_pending_old_ip_does_not_starve_new_episode(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
old_name = "episode_old"
first, _runner, _old_remote, destination = self._local_sync(
root, old_name
)
first.delete_outcomes = ["failure"]
self.assertEqual(first.run_once(), 1)
new_remote = root / "new_remote_ready"
new_remote.mkdir()
new_name = "episode_new"
self._episode(new_remote, new_name)
new_runner = LocalRsyncRunner(new_remote)
migrated = LocalSync(
remote_root=new_remote,
destination=destination,
runner=new_runner,
remote="nvidia@new-ip",
)
self.assertEqual(migrated.run_once(), 1)
self.assertFalse((new_remote / new_name).exists())
self.assertTrue((destination / new_name).is_dir())
status = json.loads((destination / "sync_status.json").read_text())
self.assertEqual(status["state"], "delete_pending")
self.assertEqual(status["pending_remote_cleanup"], 1)
self.assertIn("remote_identity", status["last_delete_error"])
def test_delete_helper_uses_independent_long_timeout(self) -> None:
class AckRunner:
def __init__(self) -> None:
self.timeouts: list[float] = []
def run(
self, command: list[str], *, timeout: float
) -> subprocess.CompletedProcess[str]:
self.timeouts.append(timeout)
payload = {
"state": "already_absent",
"episode_id": "episode_timeout",
"manifest_sha256": "0" * 64,
}
return subprocess.CompletedProcess(
command, 0, json.dumps(payload) + "\n", ""
)
with tempfile.TemporaryDirectory() as temporary:
runner = AckRunner()
syncer = DataGetSync(
remote="nvidia@test",
remote_ready=FIXED_REMOTE_READY,
destination=Path(temporary),
status_file=Path(temporary) / "status.json",
runner=runner,
delete_remote_after_sync=True,
remote_delete_timeout_s=600.0,
)
syncer._invoke_delete_helper("episode_timeout", "0" * 64)
self.assertEqual(runner.timeouts, [600.0])
def test_existing_final_without_receipt_is_deep_verified_before_delete(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
root = Path(temporary)
syncer, _runner, remote, destination = self._local_sync(root, name)
shutil.copytree(remote / name, destination / name)
with mock.patch(
"data_get_sync.validate_episode_dir",
wraps=validate_episode_dir,
) as deep_validator:
self.assertEqual(syncer.run_once(), 0)
self.assertGreaterEqual(deep_validator.call_count, 1)
self.assertTrue((destination / name / VERIFIED_RECEIPT_NAME).is_file())
self.assertFalse((remote / name).exists())
def test_existing_corrupt_final_without_receipt_never_deletes_remote(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
root = Path(temporary)
syncer, _runner, remote, destination = self._local_sync(root, name)
shutil.copytree(remote / name, destination / name)
(destination / name / "bag" / "bag_0.mcap").write_bytes(b"corrupt")
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(syncer.delete_calls, [])
self.assertTrue((remote / name).exists())
def test_existing_manifest_collision_never_deletes_remote(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
name = "episode_20260810T120000000_deadbeef"
root = Path(temporary)
syncer, _runner, remote, destination = self._local_sync(root, name)
self._episode(destination, name)
local_manifest = destination / name / "manifest.json"
local_manifest.write_bytes(local_manifest.read_bytes() + b"\n")
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(syncer.delete_calls, [])
self.assertTrue((remote / name).exists())
def test_receipt_write_failure_never_calls_delete(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
syncer, _runner, remote, _destination = self._local_sync(Path(temporary))
with mock.patch(
"data_get_sync.atomic_write_json",
side_effect=RuntimeError("injected receipt fsync failure"),
):
with self.assertRaises(RuntimeError):
syncer.sync_episode(next(remote.iterdir()).name)
self.assertEqual(syncer.delete_calls, [])
self.assertEqual(len(list(remote.iterdir())), 1)
def test_sync_flock_rejects_concurrent_once(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
first, runner, remote, destination = self._local_sync(root)
second = LocalSync(
remote_root=remote,
destination=destination,
runner=runner,
)
with first.process_lock():
with self.assertRaisesRegex(RuntimeError, "already holds"):
second.run_once()
def test_symlink_lock_is_refused(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
syncer, _runner, _remote, destination = self._local_sync(root)
sentinel = root / "sentinel"
sentinel.write_text("do not touch", encoding="utf-8")
(destination / ".data_get_sync.lock").symlink_to(sentinel)
with self.assertRaises(RuntimeError):
syncer.run_once()
self.assertEqual(sentinel.read_text(encoding="utf-8"), "do not touch")
def test_delete_mode_rejects_nonfixed_remote_ready(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
with self.assertRaises(ValueError):
DataGetSync(
remote="nvidia@test",
remote_ready="/tmp/not-ready",
destination=root,
status_file=root / "status.json",
delete_remote_after_sync=True,
)
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()