#!/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()