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